[model-gateway] introduce request ctx for oai router (#14434)

Co-authored-by: key4ng <rukeyang@gmail.com>
This commit is contained in:
Simo Lin
2025-12-04 08:31:44 -08:00
committed by GitHub
parent 788628b56f
commit b01fc161eb
5 changed files with 492 additions and 468 deletions

View File

@@ -0,0 +1,243 @@
//! Request context types for OpenAI router pipeline.
use std::sync::Arc;
use axum::http::HeaderMap;
use serde_json::Value;
use super::provider::Provider;
use crate::{
core::Worker,
data_connector::{ConversationItemStorage, ConversationStorage, ResponseStorage},
mcp::McpManager,
protocols::{chat::ChatCompletionRequest, responses::ResponsesRequest},
};
pub struct RequestContext {
pub input: RequestInput,
pub components: ComponentRefs,
pub state: ProcessingState,
}
pub struct RequestInput {
pub request_type: RequestType,
pub headers: Option<HeaderMap>,
#[allow(dead_code)]
pub model_id: Option<String>,
}
pub enum RequestType {
Chat(Arc<ChatCompletionRequest>),
Responses(Arc<ResponsesRequest>),
}
#[derive(Clone)]
pub struct SharedComponents {
pub client: reqwest::Client,
}
pub struct ResponsesComponents {
pub shared: SharedComponents,
pub mcp_manager: Arc<McpManager>,
pub response_storage: Arc<dyn ResponseStorage>,
pub conversation_storage: Arc<dyn ConversationStorage>,
pub conversation_item_storage: Arc<dyn ConversationItemStorage>,
}
pub enum ComponentRefs {
Shared(Arc<SharedComponents>),
Responses(Arc<ResponsesComponents>),
}
impl ComponentRefs {
pub fn client(&self) -> &reqwest::Client {
match self {
ComponentRefs::Shared(s) => &s.client,
ComponentRefs::Responses(r) => &r.shared.client,
}
}
pub fn mcp_manager(&self) -> Option<&Arc<McpManager>> {
match self {
ComponentRefs::Shared(_) => None,
ComponentRefs::Responses(r) => Some(&r.mcp_manager),
}
}
pub fn response_storage(&self) -> Option<&Arc<dyn ResponseStorage>> {
match self {
ComponentRefs::Shared(_) => None,
ComponentRefs::Responses(r) => Some(&r.response_storage),
}
}
pub fn conversation_storage(&self) -> Option<&Arc<dyn ConversationStorage>> {
match self {
ComponentRefs::Shared(_) => None,
ComponentRefs::Responses(r) => Some(&r.conversation_storage),
}
}
pub fn conversation_item_storage(&self) -> Option<&Arc<dyn ConversationItemStorage>> {
match self {
ComponentRefs::Shared(_) => None,
ComponentRefs::Responses(r) => Some(&r.conversation_item_storage),
}
}
}
#[derive(Default)]
pub struct ProcessingState {
pub worker: Option<WorkerSelection>,
pub payload: Option<PayloadState>,
}
pub struct WorkerSelection {
pub worker: Arc<dyn Worker>,
#[allow(dead_code)]
pub provider: Arc<dyn Provider>,
}
pub struct PayloadState {
pub json: Value,
pub url: String,
pub previous_response_id: Option<String>,
}
impl RequestContext {
pub fn for_responses(
request: Arc<ResponsesRequest>,
headers: Option<HeaderMap>,
model_id: Option<String>,
components: ComponentRefs,
) -> Self {
Self {
input: RequestInput {
request_type: RequestType::Responses(request),
headers,
model_id,
},
components,
state: ProcessingState::default(),
}
}
pub fn for_chat(
request: Arc<ChatCompletionRequest>,
headers: Option<HeaderMap>,
model_id: Option<String>,
components: ComponentRefs,
) -> Self {
Self {
input: RequestInput {
request_type: RequestType::Chat(request),
headers,
model_id,
},
components,
state: ProcessingState::default(),
}
}
}
impl RequestContext {
pub fn responses_request(&self) -> &ResponsesRequest {
match &self.input.request_type {
RequestType::Responses(req) => req.as_ref(),
_ => panic!("Expected responses request"),
}
}
#[allow(dead_code)]
pub fn responses_request_arc(&self) -> Arc<ResponsesRequest> {
match &self.input.request_type {
RequestType::Responses(req) => Arc::clone(req),
_ => panic!("Expected responses request"),
}
}
pub fn is_streaming(&self) -> bool {
match &self.input.request_type {
RequestType::Chat(req) => req.stream,
RequestType::Responses(req) => req.stream.unwrap_or(false),
}
}
pub fn headers(&self) -> Option<&HeaderMap> {
self.input.headers.as_ref()
}
#[allow(dead_code)]
pub fn model_id(&self) -> Option<&str> {
self.input.model_id.as_deref()
}
pub fn worker(&self) -> Option<&Arc<dyn Worker>> {
self.state.worker.as_ref().map(|w| &w.worker)
}
#[allow(dead_code)]
pub fn provider(&self) -> Option<&dyn Provider> {
self.state.worker.as_ref().map(|w| w.provider.as_ref())
}
pub fn payload(&self) -> Option<&PayloadState> {
self.state.payload.as_ref()
}
pub fn take_payload(&mut self) -> Option<PayloadState> {
self.state.payload.take()
}
}
pub struct StorageHandles {
pub response: Arc<dyn ResponseStorage>,
pub conversation: Arc<dyn ConversationStorage>,
pub conversation_item: Arc<dyn ConversationItemStorage>,
}
pub struct OwnedStreamingContext {
pub url: String,
pub payload: Value,
pub original_body: ResponsesRequest,
pub previous_response_id: Option<String>,
pub storage: StorageHandles,
}
impl RequestContext {
pub fn into_streaming_context(mut self) -> OwnedStreamingContext {
let payload_state = self.take_payload().expect("Payload not prepared");
OwnedStreamingContext {
url: payload_state.url,
payload: payload_state.json,
original_body: self.responses_request().clone(),
previous_response_id: payload_state.previous_response_id,
storage: StorageHandles {
response: self
.components
.response_storage()
.expect("Response storage required")
.clone(),
conversation: self
.components
.conversation_storage()
.expect("Conversation storage required")
.clone(),
conversation_item: self
.components
.conversation_item_storage()
.expect("Conversation item storage required")
.clone(),
},
}
}
}
pub struct StreamingEventContext<'a> {
pub server_label: &'a str,
pub original_request: &'a ResponsesRequest,
pub previous_response_id: Option<&'a str>,
}
pub type StreamingRequest = OwnedStreamingContext;

View File

@@ -7,6 +7,7 @@
//! - Multi-turn tool execution loops
//! - SSE (Server-Sent Events) streaming
mod context;
pub mod conversations;
pub mod mcp;
pub mod provider;

View File

@@ -225,6 +225,13 @@ impl ProviderRegistry {
.unwrap_or(self.default_provider.as_ref())
}
pub fn get_arc(&self, provider_type: &ProviderType) -> Arc<dyn Provider> {
self.providers
.get(provider_type)
.cloned()
.unwrap_or_else(|| Arc::clone(&self.default_provider))
}
pub fn get_for_model(&self, model_name: &str) -> &dyn Provider {
match ProviderType::from_model_name(model_name) {
Some(pt) => self.get(&pt),
@@ -235,122 +242,8 @@ impl ProviderRegistry {
pub fn default_provider(&self) -> &dyn Provider {
self.default_provider.as_ref()
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
#[test]
fn test_sglang_provider_passthrough() {
let provider = SGLangProvider;
let mut payload = json!({"regex": ".*", "top_k": 50});
provider
.transform_request(&mut payload, Endpoint::Chat)
.unwrap();
assert!(payload.get("regex").is_some());
assert!(payload.get("top_k").is_some());
}
#[test]
fn test_openai_provider_strips_sglang_fields() {
let provider = OpenAIProvider;
let mut payload = json!({"regex": ".*", "top_k": 50, "temperature": 0.7});
provider
.transform_request(&mut payload, Endpoint::Chat)
.unwrap();
assert!(payload.get("regex").is_none());
assert!(payload.get("top_k").is_none());
assert!(payload.get("temperature").is_some());
}
#[test]
fn test_xai_provider_transforms_responses_input() {
let provider = XAIProvider;
let mut payload = json!({
"input": [{
"id": "msg_123",
"status": "completed",
"content": [{"type": "output_text", "text": "Hello"}]
}]
});
provider
.transform_request(&mut payload, Endpoint::Responses)
.unwrap();
let item = &payload["input"][0];
assert!(item.get("id").is_none());
assert!(item.get("status").is_none());
assert_eq!(item["content"][0]["type"], "input_text");
}
#[test]
fn test_gemini_provider_removes_false_logprobs() {
let provider = GeminiProvider;
let mut payload = json!({"logprobs": false});
provider
.transform_request(&mut payload, Endpoint::Chat)
.unwrap();
assert!(payload.get("logprobs").is_none());
}
#[test]
fn test_gemini_provider_keeps_true_logprobs() {
let provider = GeminiProvider;
let mut payload = json!({"logprobs": true});
provider
.transform_request(&mut payload, Endpoint::Chat)
.unwrap();
assert_eq!(payload.get("logprobs").unwrap(), true);
}
#[test]
fn test_provider_registry_lookup() {
let registry = ProviderRegistry::new();
assert_eq!(
registry.get(&ProviderType::OpenAI).provider_type(),
ProviderType::OpenAI
);
assert_eq!(
registry.get(&ProviderType::XAI).provider_type(),
ProviderType::XAI
);
let custom = ProviderType::Custom("unknown".to_string());
assert_eq!(registry.get(&custom).provider_type(), ProviderType::OpenAI);
}
#[test]
fn test_provider_registry_get_for_model() {
let registry = ProviderRegistry::new();
assert_eq!(
registry.get_for_model("gpt-4").provider_type(),
ProviderType::OpenAI
);
assert_eq!(
registry.get_for_model("grok-2").provider_type(),
ProviderType::XAI
);
assert_eq!(
registry.get_for_model("gemini-pro").provider_type(),
ProviderType::Gemini
);
assert_eq!(
registry.get_for_model("llama-3.1-8b").provider_type(),
ProviderType::OpenAI
);
pub fn default_provider_arc(&self) -> Arc<dyn Provider> {
Arc::clone(&self.default_provider)
}
}

View File

@@ -1,5 +1,3 @@
//! OpenAI router - main coordinator that delegates to specialized modules
use std::{
any::Any,
collections::HashSet,
@@ -19,13 +17,16 @@ use tokio::sync::mpsc;
use tokio_stream::wrappers::UnboundedReceiverStream;
use tracing::warn;
// Import from sibling modules
use super::conversations::{
create_conversation, create_conversation_items, delete_conversation, delete_conversation_item,
get_conversation, get_conversation_item, list_conversation_items, persist_conversation_items,
update_conversation,
};
use super::{
context::{
ComponentRefs, PayloadState, RequestContext, ResponsesComponents, SharedComponents,
WorkerSelection,
},
conversations::{
create_conversation, create_conversation_items, delete_conversation,
delete_conversation_item, get_conversation, get_conversation_item, list_conversation_items,
persist_conversation_items, update_conversation,
},
mcp::{
ensure_request_mcp_client, execute_tool_loop, prepare_mcp_payload_for_streaming,
McpLoopConfig,
@@ -38,11 +39,7 @@ use super::{
use crate::{
app_context::AppContext,
core::{model_type::Endpoint, ModelCard, ProviderType, RuntimeType, Worker, WorkerRegistry},
data_connector::{
ConversationId, ConversationItemStorage, ConversationStorage, ListParams, ResponseId,
ResponseStorage, SortOrder,
},
mcp::McpManager,
data_connector::{ConversationId, ListParams, ResponseId, SortOrder},
protocols::{
chat::ChatCompletionRequest,
classify::ClassifyRequest,
@@ -57,32 +54,12 @@ use crate::{
},
};
// ============================================================================
// OpenAIRouter Struct
// ============================================================================
/// Router for OpenAI backend
///
/// This router manages connections to OpenAI-compatible API endpoints (OpenAI, xAI, etc.)
/// using the Worker abstraction. Workers are registered via the external worker registration
/// workflow and stored in the WorkerRegistry.
pub struct OpenAIRouter {
/// HTTP client for upstream OpenAI-compatible API
client: reqwest::Client,
/// Worker registry for model-based worker lookup
worker_registry: Arc<WorkerRegistry>,
/// Provider registry for vendor-specific transformations
provider_registry: ProviderRegistry,
/// Health status
healthy: AtomicBool,
/// Response storage for managing conversation history
response_storage: Arc<dyn ResponseStorage>,
/// Conversation storage backend
conversation_storage: Arc<dyn ConversationStorage>,
/// Conversation item storage backend
conversation_item_storage: Arc<dyn ConversationItemStorage>,
/// MCP manager (handles both static and dynamic servers)
mcp_manager: Arc<McpManager>,
shared_components: Arc<SharedComponents>,
responses_components: Arc<ResponsesComponents>,
}
impl std::fmt::Debug for OpenAIRouter {
@@ -98,76 +75,70 @@ impl std::fmt::Debug for OpenAIRouter {
}
impl OpenAIRouter {
/// Maximum number of conversation items to attach as input when a conversation is provided
const MAX_CONVERSATION_HISTORY_ITEMS: usize = 100;
/// Create a new OpenAI router
///
/// Workers are registered separately via the external worker registration workflow.
/// This router queries the WorkerRegistry to find workers that support requested models.
fn shared_components(&self) -> Arc<SharedComponents> {
Arc::clone(&self.shared_components)
}
fn responses_components(&self) -> Arc<ResponsesComponents> {
Arc::clone(&self.responses_components)
}
pub async fn new(ctx: &Arc<AppContext>) -> Result<Self, String> {
// Use HTTP client from AppContext
let client = ctx.client.clone();
// Get worker registry from AppContext
let worker_registry = ctx.worker_registry.clone();
// Get MCP manager from AppContext (must be initialized)
let mcp_manager = ctx
.mcp_manager
.get()
.ok_or_else(|| "MCP manager not initialized in AppContext".to_string())?
.clone();
Ok(Self {
client,
worker_registry,
provider_registry: ProviderRegistry::new(),
healthy: AtomicBool::new(true),
let shared_components = Arc::new(SharedComponents {
client: ctx.client.clone(),
});
let responses_components = Arc::new(ResponsesComponents {
shared: SharedComponents {
client: ctx.client.clone(),
},
mcp_manager: mcp_manager.clone(),
response_storage: ctx.response_storage.clone(),
conversation_storage: ctx.conversation_storage.clone(),
conversation_item_storage: ctx.conversation_item_storage.clone(),
mcp_manager,
});
Ok(Self {
worker_registry,
provider_registry: ProviderRegistry::new(),
healthy: AtomicBool::new(true),
shared_components,
responses_components,
})
}
/// Get the provider for a worker and optional model.
///
/// Priority:
/// 1. Worker's provider for the specific model (if worker knows about it)
/// 2. Infer from model name (ProviderType::from_model_name)
/// 3. Default provider (SGLang passthrough)
fn get_provider_for_worker<'a>(
&'a self,
fn get_provider_arc_for_worker(
&self,
worker: &dyn Worker,
model_id: Option<&str>,
) -> &'a dyn super::provider::Provider {
// Try worker's provider for the model first
) -> Arc<dyn super::provider::Provider> {
if let Some(model) = model_id {
if let Some(pt) = worker.provider_for_model(model) {
return self.provider_registry.get(pt);
return self.provider_registry.get_arc(pt);
}
// Fall back to model name inference
if let Some(pt) = ProviderType::from_model_name(model) {
return self.provider_registry.get(&pt);
return self.provider_registry.get_arc(&pt);
}
}
// Default to SGLang passthrough
self.provider_registry.default_provider()
self.provider_registry.default_provider_arc()
}
/// Refresh models for a single external worker by querying its /v1/models endpoint.
///
/// Returns true if refresh succeeded and models were cached on the worker.
async fn refresh_worker_models(
&self,
worker: &Arc<dyn Worker>,
auth_header: Option<&HeaderValue>,
) -> bool {
let url = format!("{}/v1/models", worker.url());
// Build request to backend
let mut backend_req = self.client.get(&url);
let mut backend_req = self.shared_components.client.get(&url);
if let Some(auth) = auth_header {
backend_req = apply_provider_headers(backend_req, &url, Some(auth));
}
@@ -216,7 +187,6 @@ impl OpenAIRouter {
}
}
/// Refresh models for ALL external workers in parallel.
async fn refresh_external_models(&self, auth_header: Option<&HeaderValue>) {
let external_workers = self.worker_registry.get_workers_filtered(
None,
@@ -235,7 +205,6 @@ impl OpenAIRouter {
external_workers.len()
);
// Refresh all workers in parallel
let futures: Vec<_> = external_workers
.iter()
.map(|w| self.refresh_worker_models(w, auth_header))
@@ -244,41 +213,19 @@ impl OpenAIRouter {
join_all(futures).await;
}
/// Select a worker for the given model using the WorkerRegistry.
///
/// This method queries the registry for external workers (RuntimeType::External)
/// that support the requested model. It checks:
/// 1. Workers registered with matching model ID (including aliases via ModelCard)
/// 2. Worker health status
/// 3. Circuit breaker state
///
/// If no worker is found with explicit model support, it will refresh models
/// on all external workers in parallel, then retry the search.
///
/// Returns an error response if no suitable worker is found.
async fn select_worker_for_model(
&self,
model_id: &str,
auth_header: Option<&HeaderValue>,
) -> Result<Arc<dyn Worker>, Box<Response>> {
// Helper to find candidates for a model
// Note: We get ALL external workers and filter by supports_model() because
// wildcard workers (empty models) aren't in the model index but support any model
let find_candidates = || {
self.worker_registry
.get_workers_filtered(
None, // Get all external workers, not just those in model index
None,
None,
Some(RuntimeType::External),
true, // healthy_only
)
.get_workers_filtered(None, None, None, Some(RuntimeType::External), true)
.into_iter()
.filter(|w| w.supports_model(model_id) && w.circuit_breaker().can_execute())
.collect::<Vec<_>>()
};
// First try: find workers that already support this model
let candidates = find_candidates();
if !candidates.is_empty() {
return Ok(candidates
@@ -287,14 +234,12 @@ impl OpenAIRouter {
.expect("candidates is not empty"));
}
// No match found - refresh models on all external workers
tracing::debug!(
"No worker found for model '{}', refreshing external worker models",
model_id
);
self.refresh_external_models(auth_header).await;
// Second try: check if any worker now supports the model after refresh
let candidates = find_candidates();
if !candidates.is_empty() {
return Ok(candidates
@@ -317,43 +262,35 @@ impl OpenAIRouter {
))
}
/// Handle non-streaming response with optional MCP tool loop
async fn handle_non_streaming_response(
&self,
worker: &Arc<dyn Worker>,
headers: Option<&HeaderMap>,
mut payload: Value,
original_body: &ResponsesRequest,
original_previous_response_id: Option<String>,
) -> Response {
let url = format!("{}/v1/responses", worker.url());
async fn handle_non_streaming_response(&self, mut ctx: RequestContext) -> Response {
let payload_state = ctx.take_payload().expect("Payload not prepared");
let mut payload = payload_state.json;
let url = payload_state.url;
let previous_response_id = payload_state.previous_response_id;
let original_body = ctx.responses_request();
let worker = ctx.worker().expect("Worker not selected");
let mcp_manager = ctx.components.mcp_manager().expect("MCP manager required");
// Check if MCP is active for this request
// Ensure dynamic client is created if needed
if let Some(ref tools) = original_body.tools {
ensure_request_mcp_client(&self.mcp_manager, tools.as_slice()).await;
ensure_request_mcp_client(mcp_manager, tools.as_slice()).await;
}
// Use the tool loop if the manager has any tools available (static or dynamic).
let active_mcp = if self.mcp_manager.list_tools().is_empty() {
let active_mcp = if mcp_manager.list_tools().is_empty() {
None
} else {
Some(&self.mcp_manager)
Some(mcp_manager)
};
let mut response_json: Value;
// If MCP is active, execute tool loop
if let Some(mcp) = active_mcp {
let config = McpLoopConfig::default();
// Transform MCP tools to function tools
prepare_mcp_payload_for_streaming(&mut payload, mcp);
match execute_tool_loop(
&self.client,
ctx.components.client(),
&url,
headers,
ctx.headers(),
payload,
original_body,
mcp,
@@ -372,13 +309,8 @@ impl OpenAIRouter {
}
}
} else {
// No MCP - simple request
let mut request_builder = self.client.post(&url).json(&payload);
// Apply provider-specific headers (handles Anthropic x-api-key, etc.)
// Passthrough mode: user's auth header takes priority, worker's key is fallback
let auth_header = extract_auth_header(headers, worker.api_key());
let mut request_builder = ctx.components.client().post(&url).json(&payload);
let auth_header = extract_auth_header(ctx.headers(), worker.api_key());
request_builder = apply_provider_headers(request_builder, &url, auth_header.as_ref());
let response = match request_builder.send().await {
@@ -421,19 +353,26 @@ impl OpenAIRouter {
worker.circuit_breaker().record_success();
}
// Patch response with metadata
mask_tools_as_mcp(&mut response_json, original_body);
patch_streaming_response_json(
&mut response_json,
original_body,
original_previous_response_id.as_deref(),
previous_response_id.as_deref(),
);
// Always persist conversation items and response (even without conversation)
if let Err(err) = persist_conversation_items(
self.conversation_storage.clone(),
self.conversation_item_storage.clone(),
self.response_storage.clone(),
ctx.components
.conversation_storage()
.expect("Conversation storage required")
.clone(),
ctx.components
.conversation_item_storage()
.expect("Conversation item storage required")
.clone(),
ctx.components
.response_storage()
.expect("Response storage required")
.clone(),
&response_json,
original_body,
)
@@ -446,10 +385,6 @@ impl OpenAIRouter {
}
}
// ============================================================================
// RouterTrait Implementation
// ============================================================================
#[async_trait::async_trait]
impl crate::routers::RouterTrait for OpenAIRouter {
fn as_any(&self) -> &dyn Any {
@@ -457,7 +392,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
async fn health_generate(&self, _req: Request<Body>) -> Response {
// Check health of all external workers
let external_workers: Vec<_> = self
.worker_registry
.get_all()
@@ -530,7 +464,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
async fn get_models(&self, req: Request<Body>) -> Response {
// Return models from all registered external workers
let external_workers: Vec<_> = self
.worker_registry
.get_all()
@@ -546,11 +479,9 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.into_response();
}
// Refresh models for all external workers using user's auth header
let auth_header = extract_auth_header(Some(req.headers()), &None);
self.refresh_external_models(auth_header.as_ref()).await;
// Collect models from all workers
let mut all_models = Vec::new();
let mut seen_models = HashSet::new();
@@ -562,7 +493,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.map(|p| format!("{:?}", p).to_lowercase())
.unwrap_or_else(|| "unknown".to_string());
// Add primary model ID
if seen_models.insert(model_card.id.clone()) {
all_models.push(json!({
"id": &model_card.id,
@@ -574,7 +504,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}));
}
// Add aliases as separate entries for compatibility
for alias in &model_card.aliases {
if seen_models.insert(alias.clone()) {
all_models.push(json!({
@@ -589,7 +518,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
}
// Return aggregated models
let response_json = json!({
"object": "list",
"data": all_models
@@ -599,7 +527,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
async fn get_model_info(&self, _req: Request<Body>) -> Response {
// Not directly supported without model param; return 501
(
StatusCode::NOT_IMPLEMENTED,
"get_model_info not implemented for OpenAI router",
@@ -613,7 +540,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
_body: &GenerateRequest,
_model_id: Option<&str>,
) -> Response {
// Generate endpoint is SGLang-specific, not supported for OpenAI backend
(
StatusCode::NOT_IMPLEMENTED,
"Generate endpoint not supported for OpenAI backend",
@@ -627,10 +553,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
body: &ChatCompletionRequest,
model_id: Option<&str>,
) -> Response {
// Extract auth header for passthrough mode
let auth_header = extract_auth_header(headers, &None);
// Select worker for model (discovery happens inside if needed)
let worker = match self
.select_worker_for_model(body.model.as_str(), auth_header.as_ref())
.await
@@ -639,7 +563,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
Err(response) => return *response,
};
// Serialize request body, removing SGLang-only fields
let mut payload = match to_value(body) {
Ok(v) => v,
Err(e) => {
@@ -650,8 +573,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.into_response();
}
};
// Apply provider-specific transformations
let provider = self.get_provider_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) {
return (
StatusCode::BAD_REQUEST,
@@ -660,16 +583,31 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.into_response();
}
let url = format!("{}/v1/chat/completions", worker.url());
let mut req = self.client.post(&url).json(&payload);
let mut ctx = RequestContext::for_chat(
Arc::new(body.clone()),
headers.cloned(),
model_id.map(String::from),
ComponentRefs::Shared(self.shared_components()),
);
// Apply provider-specific headers (handles Anthropic x-api-key, etc.)
// Passthrough mode: user's auth header takes priority, worker's key is fallback
let auth_header = extract_auth_header(headers, worker.api_key());
ctx.state.worker = Some(WorkerSelection {
worker: Arc::clone(&worker),
provider,
});
let url = format!("{}/v1/chat/completions", worker.url());
ctx.state.payload = Some(PayloadState {
json: payload,
url: url.clone(),
previous_response_id: None,
});
let payload_ref = ctx.payload().expect("Payload not prepared");
let mut req = ctx.components.client().post(&url).json(&payload_ref.json);
let auth_header = extract_auth_header(ctx.headers(), worker.api_key());
req = apply_provider_headers(req, &url, auth_header.as_ref());
// Accept SSE when stream=true
if body.stream {
if ctx.is_streaming() {
req = req.header("Accept", "text/event-stream");
}
@@ -688,8 +626,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
let status = StatusCode::from_u16(resp.status().as_u16())
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
if !body.stream {
// Capture Content-Type before consuming response body
if !ctx.is_streaming() {
let content_type = resp.headers().get(CONTENT_TYPE).cloned();
match resp.bytes().await {
Ok(body) => {
@@ -711,7 +648,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
}
} else {
// Stream SSE bytes to client
let stream = resp.bytes_stream();
let (tx, rx) = mpsc::unbounded_channel();
tokio::spawn(async move {
@@ -745,10 +681,9 @@ impl crate::routers::RouterTrait for OpenAIRouter {
_body: &CompletionRequest,
_model_id: Option<&str>,
) -> Response {
// Completion endpoint not implemented for OpenAI backend
(
StatusCode::NOT_IMPLEMENTED,
"Completion endpoint not implemented for OpenAI backend",
"Completion endpoint not implemented",
)
.into_response()
}
@@ -759,10 +694,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
body: &ResponsesRequest,
model_id: Option<&str>,
) -> Response {
// Extract auth header for passthrough mode
let auth_header = extract_auth_header(headers, &None);
// Select worker for model (discovery happens inside if needed)
let model = model_id.unwrap_or(body.model.as_str());
let worker = match self
.select_worker_for_model(model, auth_header.as_ref())
@@ -772,22 +705,19 @@ impl crate::routers::RouterTrait for OpenAIRouter {
Err(response) => return *response,
};
// Clone the body for validation and logic, but we'll build payload differently
let mut request_body = body.clone();
if let Some(model) = model_id {
request_body.model = model.to_string();
}
// Do not forward conversation field upstream; retain for local persistence only
request_body.conversation = None;
// Store the original previous_response_id for the response
let original_previous_response_id = request_body.previous_response_id.clone();
// Handle previous_response_id by loading prior context
let mut conversation_items: Option<Vec<ResponseInputOutputItem>> = None;
if let Some(prev_id_str) = request_body.previous_response_id.clone() {
let prev_id = ResponseId::from(prev_id_str.as_str());
match self
.responses_components
.response_storage
.get_response_chain(&prev_id, None)
.await
@@ -795,7 +725,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
Ok(chain) => {
let mut items = Vec::new();
for stored in chain.responses.iter() {
// Convert input items from stored input (which is now a JSON array)
if let Some(input_arr) = stored.input.as_array() {
for item in input_arr {
match serde_json::from_value::<ResponseInputOutputItem>(
@@ -814,7 +743,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
}
// Convert output items from stored output (which is now a JSON array)
if let Some(output_arr) = stored.output.as_array() {
for item in output_arr {
match serde_json::from_value::<ResponseInputOutputItem>(
@@ -842,12 +770,15 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
}
// Handle conversation by loading history
if let Some(conv_id_str) = body.conversation.clone() {
let conv_id = ConversationId::from(conv_id_str.as_str());
// Verify conversation exists
if let Ok(None) = self.conversation_storage.get_conversation(&conv_id).await {
if let Ok(None) = self
.responses_components
.conversation_storage
.get_conversation(&conv_id)
.await
{
return (
StatusCode::NOT_FOUND,
Json(json!({"error": "Conversation not found"})),
@@ -855,7 +786,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.into_response();
}
// Load conversation history (ascending order for chronological context)
let params = ListParams {
limit: Self::MAX_CONVERSATION_HISTORY_ITEMS,
order: SortOrder::Asc,
@@ -863,6 +793,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
};
match self
.responses_components
.conversation_item_storage
.list_items(&conv_id, params)
.await
@@ -870,8 +801,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
Ok(stored_items) => {
let mut items: Vec<ResponseInputOutputItem> = Vec::new();
for item in stored_items.into_iter() {
// Include messages, function calls, and function call outputs
// Skip reasoning items as they're internal processing details
match item.item_type.as_str() {
"message" => {
match serde_json::from_value::<Vec<ResponseContentPart>>(
@@ -897,7 +826,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
}
"function_call" => {
// The entire function_call item is stored in content field
match serde_json::from_value::<ResponseInputOutputItem>(
item.content.clone(),
) {
@@ -911,7 +839,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
}
"function_call_output" => {
// The entire function_call_output item is stored in content field
tracing::debug!(
"Loading function_call_output from DB - content: {}",
serde_json::to_string_pretty(&item.content)
@@ -934,17 +861,13 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
}
}
"reasoning" => {
// Skip reasoning items - they're internal processing details
}
"reasoning" => {}
_ => {
// Skip unknown item types
warn!("Unknown item type in conversation: {}", item.item_type);
}
}
}
// Append current request
match &request_body.input {
ResponseInput::Text(text) => {
items.push(ResponseInputOutputItem::Message {
@@ -957,7 +880,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
});
}
ResponseInput::Items(current_items) => {
// Process all item types, converting SimpleInputMessage to Message
for item in current_items.iter() {
let normalized =
crate::protocols::responses::normalize_input_item(item);
@@ -974,9 +896,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
}
// If we have conversation_items from previous_response_id, use them
if let Some(mut items) = conversation_items {
// Append current request
match &request_body.input {
ResponseInput::Text(text) => {
items.push(ResponseInputOutputItem::Message {
@@ -992,7 +912,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
});
}
ResponseInput::Items(current_items) => {
// Process all item types, converting SimpleInputMessage to Message
for item in current_items.iter() {
let normalized = crate::protocols::responses::normalize_input_item(item);
items.push(normalized);
@@ -1003,14 +922,11 @@ impl crate::routers::RouterTrait for OpenAIRouter {
request_body.input = ResponseInput::Items(items);
}
// Always set store=false for upstream (we store internally)
request_body.store = Some(false);
// Filter out reasoning items from input - they're internal processing details
if let ResponseInput::Items(ref mut items) = request_body.input {
items.retain(|item| !matches!(item, ResponseInputOutputItem::Reasoning { .. }));
}
// Convert to JSON and strip SGLang-specific fields
let mut payload = match to_value(&request_body) {
Ok(v) => v,
Err(e) => {
@@ -1022,8 +938,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
};
// Apply provider-specific transformations (handles SGLang fields, XAI/Grok, etc.)
let provider = self.get_provider_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) {
return (
StatusCode::BAD_REQUEST,
@@ -1032,32 +947,28 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.into_response();
}
// Delegate to streaming or non-streaming handler
let url = format!("{}/v1/responses", worker.url());
if body.stream.unwrap_or(false) {
handle_streaming_response(
&self.client,
worker.circuit_breaker(),
Some(&self.mcp_manager),
self.response_storage.clone(),
self.conversation_storage.clone(),
self.conversation_item_storage.clone(),
url,
headers,
payload,
body,
original_previous_response_id,
)
.await
let mut ctx = RequestContext::for_responses(
Arc::new(body.clone()),
headers.cloned(),
model_id.map(String::from),
ComponentRefs::Responses(self.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,
});
if ctx.is_streaming() {
handle_streaming_response(ctx).await
} else {
self.handle_non_streaming_response(
&worker,
headers,
payload,
body,
original_previous_response_id,
)
.await
self.handle_non_streaming_response(ctx).await
}
}
@@ -1068,7 +979,12 @@ impl crate::routers::RouterTrait for OpenAIRouter {
_params: &ResponsesGetParams,
) -> Response {
let id = ResponseId::from(response_id);
match self.response_storage.get_response(&id).await {
match self
.responses_components
.response_storage
.get_response(&id)
.await
{
Ok(Some(stored)) => {
let mut response_json = stored.raw_response;
if let Some(obj) = response_json.as_object_mut() {
@@ -1104,20 +1020,22 @@ impl crate::routers::RouterTrait for OpenAIRouter {
) -> Response {
let resp_id = ResponseId::from(response_id);
match self.response_storage.get_response(&resp_id).await {
match self
.responses_components
.response_storage
.get_response(&resp_id)
.await
{
Ok(Some(stored)) => {
// Extract items from input field (which is a JSON array)
let items = match &stored.input {
Value::Array(arr) => arr.clone(),
_ => vec![],
};
// Generate IDs for items if they don't have them
let items_with_ids: Vec<Value> = items
.into_iter()
.map(|mut item| {
if item.get("id").is_none() {
// Generate ID if not present using centralized utility
if let Some(obj) = item.as_object_mut() {
obj.insert("id".to_string(), json!(generate_id("msg")));
}
@@ -1194,7 +1112,11 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
async fn create_conversation(&self, _headers: Option<&HeaderMap>, body: &Value) -> Response {
create_conversation(&self.conversation_storage, body.clone()).await
create_conversation(
&self.responses_components.conversation_storage,
body.clone(),
)
.await
}
async fn get_conversation(
@@ -1202,7 +1124,11 @@ impl crate::routers::RouterTrait for OpenAIRouter {
_headers: Option<&HeaderMap>,
conversation_id: &str,
) -> Response {
get_conversation(&self.conversation_storage, conversation_id).await
get_conversation(
&self.responses_components.conversation_storage,
conversation_id,
)
.await
}
async fn update_conversation(
@@ -1211,7 +1137,12 @@ impl crate::routers::RouterTrait for OpenAIRouter {
conversation_id: &str,
body: &Value,
) -> Response {
update_conversation(&self.conversation_storage, conversation_id, body.clone()).await
update_conversation(
&self.responses_components.conversation_storage,
conversation_id,
body.clone(),
)
.await
}
async fn delete_conversation(
@@ -1219,7 +1150,11 @@ impl crate::routers::RouterTrait for OpenAIRouter {
_headers: Option<&HeaderMap>,
conversation_id: &str,
) -> Response {
delete_conversation(&self.conversation_storage, conversation_id).await
delete_conversation(
&self.responses_components.conversation_storage,
conversation_id,
)
.await
}
async fn list_conversation_items(
@@ -1242,8 +1177,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
list_conversation_items(
&self.conversation_storage,
&self.conversation_item_storage,
&self.responses_components.conversation_storage,
&self.responses_components.conversation_item_storage,
conversation_id,
query_params,
)
@@ -1257,8 +1192,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
body: &Value,
) -> Response {
create_conversation_items(
&self.conversation_storage,
&self.conversation_item_storage,
&self.responses_components.conversation_storage,
&self.responses_components.conversation_item_storage,
conversation_id,
body.clone(),
)
@@ -1273,8 +1208,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
include: Option<Vec<String>>,
) -> Response {
get_conversation_item(
&self.conversation_storage,
&self.conversation_item_storage,
&self.responses_components.conversation_storage,
&self.responses_components.conversation_item_storage,
conversation_id,
item_id,
include,
@@ -1289,8 +1224,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
item_id: &str,
) -> Response {
delete_conversation_item(
&self.conversation_storage,
&self.conversation_item_storage,
&self.responses_components.conversation_storage,
&self.responses_components.conversation_item_storage,
conversation_id,
item_id,
)

View File

@@ -22,8 +22,9 @@ use tokio_stream::wrappers::UnboundedReceiverStream;
use tracing::warn;
// Import from sibling modules
use super::conversations::persist_conversation_items;
use super::context::{RequestContext, StreamingEventContext, StreamingRequest};
use super::{
conversations::persist_conversation_items,
mcp::{
build_resume_payload, ensure_request_mcp_client, execute_streaming_tool_calls,
inject_mcp_metadata_streaming, prepare_mcp_payload_for_streaming,
@@ -33,7 +34,6 @@ use super::{
utils::{event_types, FunctionCallInProgress, OutputIndexMapper, StreamAction},
};
use crate::{
data_connector::{ConversationItemStorage, ConversationStorage, ResponseStorage},
protocols::responses::{ResponseToolType, ResponsesRequest},
routers::header_utils::{apply_request_headers, preserve_response_headers},
};
@@ -550,9 +550,7 @@ pub(super) fn parse_sse_block(block: &str) -> (Option<&str>, Cow<'_, str>) {
/// Returns true if any changes were made
pub(super) fn apply_event_transformations_inplace(
parsed_data: &mut Value,
server_label: &str,
original_request: &ResponsesRequest,
previous_response_id: Option<&str>,
ctx: &StreamingEventContext<'_>,
) -> bool {
let mut changed = false;
@@ -575,13 +573,13 @@ pub(super) fn apply_event_transformations_inplace(
.get_mut("response")
.and_then(|v| v.as_object_mut())
{
let desired_store = Value::Bool(original_request.store.unwrap_or(false));
let desired_store = Value::Bool(ctx.original_request.store.unwrap_or(false));
if response_obj.get("store") != Some(&desired_store) {
response_obj.insert("store".to_string(), desired_store);
changed = true;
}
if let Some(prev_id) = previous_response_id {
if let Some(prev_id) = ctx.previous_response_id {
let needs_previous = response_obj
.get("previous_response_id")
.map(|v| v.is_null() || v.as_str().map(|s| s.is_empty()).unwrap_or(false))
@@ -598,7 +596,8 @@ pub(super) fn apply_event_transformations_inplace(
// Mask tools from function to MCP format (optimized without cloning)
if response_obj.get("tools").is_some() {
let requested_mcp = original_request
let requested_mcp = ctx
.original_request
.tools
.as_ref()
.map(|tools| {
@@ -609,7 +608,7 @@ pub(super) fn apply_event_transformations_inplace(
.unwrap_or(false);
if requested_mcp {
if let Some(mcp_tools) = build_mcp_tools_value(original_request) {
if let Some(mcp_tools) = build_mcp_tools_value(ctx.original_request) {
response_obj.insert("tools".to_string(), mcp_tools);
response_obj
.entry("tool_choice".to_string())
@@ -630,7 +629,7 @@ pub(super) fn apply_event_transformations_inplace(
|| item_type == event_types::ITEM_TYPE_FUNCTION_TOOL_CALL
{
item["type"] = json!(event_types::ITEM_TYPE_MCP_CALL);
item["server_label"] = json!(server_label);
item["server_label"] = json!(ctx.server_label);
// Transform ID from fc_* to mcp_*
if let Some(id) = item.get("id").and_then(|v| v.as_str()) {
@@ -682,16 +681,13 @@ fn build_mcp_tools_value(original_body: &ResponsesRequest) -> Option<Value> {
/// Forward and transform a streaming event to the client
/// Returns false if client disconnected
#[allow(clippy::too_many_arguments)]
pub(super) fn forward_streaming_event(
raw_block: &str,
event_name: Option<&str>,
data: &str,
handler: &mut StreamingToolHandler,
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
server_label: &str,
original_request: &ResponsesRequest,
previous_response_id: Option<&str>,
ctx: &StreamingEventContext<'_>,
sequence_number: &mut u64,
) -> bool {
// Skip individual function_call_arguments.delta events - we'll send them as one
@@ -808,12 +804,7 @@ pub(super) fn forward_streaming_event(
}
// Apply all transformations in-place (single parse/serialize!)
apply_event_transformations_inplace(
&mut parsed_data,
server_label,
original_request,
previous_response_id,
);
apply_event_transformations_inplace(&mut parsed_data, ctx);
if let Some(response_obj) = parsed_data
.get_mut("response")
@@ -899,16 +890,13 @@ pub(super) fn forward_streaming_event(
/// Send final response.completed event to client
/// Returns false if client disconnected
#[allow(clippy::too_many_arguments)]
pub(super) fn send_final_response_event(
handler: &StreamingToolHandler,
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
sequence_number: &mut u64,
state: &ToolLoopState,
active_mcp: Option<&Arc<crate::mcp::McpManager>>,
original_request: &ResponsesRequest,
previous_response_id: Option<&str>,
server_label: &str,
ctx: &StreamingEventContext<'_>,
) -> bool {
let mut final_response = match handler.snapshot_final_response() {
Some(resp) => resp,
@@ -925,11 +913,15 @@ pub(super) fn send_final_response_event(
}
if let Some(mcp) = active_mcp {
inject_mcp_metadata_streaming(&mut final_response, state, mcp, server_label);
inject_mcp_metadata_streaming(&mut final_response, state, mcp, ctx.server_label);
}
mask_tools_as_mcp(&mut final_response, original_request);
patch_streaming_response_json(&mut final_response, original_request, previous_response_id);
mask_tools_as_mcp(&mut final_response, ctx.original_request);
patch_streaming_response_json(
&mut final_response,
ctx.original_request,
ctx.previous_response_id,
);
if let Some(obj) = final_response.as_object_mut() {
obj.insert("status".to_string(), Value::String("completed".to_string()));
@@ -955,20 +947,13 @@ pub(super) fn send_final_response_event(
// ============================================================================
/// Simple pass-through streaming without MCP interception
#[allow(clippy::too_many_arguments)]
pub(super) async fn handle_simple_streaming_passthrough(
client: &reqwest::Client,
circuit_breaker: &crate::core::CircuitBreaker,
response_storage: Arc<dyn ResponseStorage>,
conversation_storage: Arc<dyn ConversationStorage>,
conversation_item_storage: Arc<dyn ConversationItemStorage>,
url: String,
headers: Option<&HeaderMap>,
payload: Value,
original_body: &ResponsesRequest,
original_previous_response_id: Option<String>,
req: StreamingRequest,
) -> Response {
let mut request_builder = client.post(&url).json(&payload);
let mut request_builder = client.post(&req.url).json(&req.payload);
if let Some(headers) = headers {
request_builder = apply_request_headers(headers, request_builder, true);
@@ -1008,10 +993,11 @@ pub(super) async fn handle_simple_streaming_passthrough(
let (tx, rx) = mpsc::unbounded_channel::<Result<Bytes, io::Error>>();
let should_store = original_body.store.unwrap_or(false);
let original_request = original_body.clone();
let should_store = req.original_body.store.unwrap_or(false);
let original_request = req.original_body;
let persist_needed = original_request.conversation.is_some();
let previous_response_id = original_previous_response_id.clone();
let previous_response_id = req.previous_response_id;
let storage = req.storage;
tokio::spawn(async move {
let mut accumulator = StreamingResponseAccumulator::new();
@@ -1090,9 +1076,9 @@ pub(super) async fn handle_simple_streaming_passthrough(
// Always persist conversation items and response (even without conversation)
if let Err(err) = persist_conversation_items(
conversation_storage.clone(),
conversation_item_storage.clone(),
response_storage.clone(),
storage.conversation.clone(),
storage.conversation_item.clone(),
storage.response.clone(),
&response_json,
&original_request,
)
@@ -1125,27 +1111,23 @@ pub(super) async fn handle_simple_streaming_passthrough(
}
/// Handle streaming WITH MCP tool call interception and execution
#[allow(clippy::too_many_arguments)]
pub(super) async fn handle_streaming_with_tool_interception(
client: &reqwest::Client,
response_storage: Arc<dyn ResponseStorage>,
conversation_storage: Arc<dyn ConversationStorage>,
conversation_item_storage: Arc<dyn ConversationItemStorage>,
url: String,
headers: Option<&HeaderMap>,
mut payload: Value,
original_body: &ResponsesRequest,
original_previous_response_id: Option<String>,
req: StreamingRequest,
active_mcp: &Arc<crate::mcp::McpManager>,
) -> Response {
// Transform MCP tools to function tools in payload
let mut payload = req.payload;
prepare_mcp_payload_for_streaming(&mut payload, active_mcp);
let (tx, rx) = mpsc::unbounded_channel::<Result<Bytes, io::Error>>();
let should_store = original_body.store.unwrap_or(false);
let original_request = original_body.clone();
let should_store = req.original_body.store.unwrap_or(false);
let original_request = req.original_body;
let persist_needed = original_request.conversation.is_some();
let previous_response_id = original_previous_response_id.clone();
let previous_response_id = req.previous_response_id;
let url = req.url;
let storage = req.storage;
let client_clone = client.clone();
let url_clone = url.clone();
@@ -1178,6 +1160,12 @@ pub(super) async fn handle_streaming_with_tool_interception(
})
.unwrap_or("mcp");
let streaming_ctx = StreamingEventContext {
server_label,
original_request: &original_request,
previous_response_id: previous_response_id.as_deref(),
};
loop {
// Make streaming request
let mut request_builder = client_clone.post(&url_clone).json(&current_payload);
@@ -1271,9 +1259,7 @@ pub(super) async fn handle_streaming_with_tool_interception(
data.as_ref(),
&mut handler,
&tx,
server_label,
&original_request,
previous_response_id.as_deref(),
&streaming_ctx,
&mut sequence_number,
) {
// Client disconnected
@@ -1319,9 +1305,7 @@ pub(super) async fn handle_streaming_with_tool_interception(
data.as_ref(),
&mut handler,
&tx,
server_label,
&original_request,
previous_response_id.as_deref(),
&streaming_ctx,
&mut sequence_number,
) {
// Client disconnected
@@ -1358,9 +1342,7 @@ pub(super) async fn handle_streaming_with_tool_interception(
&mut sequence_number,
&state,
Some(&active_mcp_clone),
&original_request,
previous_response_id.as_deref(),
server_label,
&streaming_ctx,
) {
return;
}
@@ -1393,9 +1375,9 @@ pub(super) async fn handle_streaming_with_tool_interception(
// Always persist conversation items and response (even without conversation)
if let Err(err) = persist_conversation_items(
conversation_storage.clone(),
conversation_item_storage.clone(),
response_storage.clone(),
storage.conversation.clone(),
storage.conversation_item.clone(),
storage.response.clone(),
&response_json,
&original_request,
)
@@ -1483,50 +1465,32 @@ pub(super) async fn handle_streaming_with_tool_interception(
response
}
/// Main entry point for handling streaming responses
/// Delegates to simple passthrough or MCP tool interception based on configuration
#[allow(clippy::too_many_arguments)]
pub(super) async fn handle_streaming_response(
client: &reqwest::Client,
circuit_breaker: &crate::core::CircuitBreaker,
mcp_manager: Option<&Arc<crate::mcp::McpManager>>,
response_storage: Arc<dyn ResponseStorage>,
conversation_storage: Arc<dyn ConversationStorage>,
conversation_item_storage: Arc<dyn ConversationItemStorage>,
url: String,
headers: Option<&HeaderMap>,
payload: Value,
original_body: &ResponsesRequest,
original_previous_response_id: Option<String>,
) -> Response {
// Check if MCP is active for this request
// Ensure dynamic client is created if needed
if let (Some(manager), Some(ref tools)) = (mcp_manager, &original_body.tools) {
ensure_request_mcp_client(manager, tools.as_slice()).await;
pub(super) async fn handle_streaming_response(ctx: RequestContext) -> Response {
let worker = ctx.worker().expect("Worker not selected").clone();
let circuit_breaker = worker.circuit_breaker();
let headers = ctx.headers().cloned();
let original_body = ctx.responses_request();
let mcp_manager = ctx.components.mcp_manager().expect("MCP manager required");
if let Some(ref tools) = original_body.tools {
ensure_request_mcp_client(mcp_manager, tools.as_slice()).await;
}
// Use the tool loop if the manager has any tools available (static or dynamic).
let active_mcp = mcp_manager.and_then(|mgr| {
if mgr.list_tools().is_empty() {
None
} else {
Some(mgr)
}
});
let active_mcp = if mcp_manager.list_tools().is_empty() {
None
} else {
Some(mcp_manager.clone())
};
let client = ctx.components.client().clone();
let req = ctx.into_streaming_context();
// If no MCP is active, use simple pass-through streaming
if active_mcp.is_none() {
return handle_simple_streaming_passthrough(
client,
&client,
circuit_breaker,
response_storage,
conversation_storage,
conversation_item_storage,
url,
headers,
payload,
original_body,
original_previous_response_id,
headers.as_ref(),
req,
)
.await;
}
@@ -1534,17 +1498,5 @@ pub(super) async fn handle_streaming_response(
let active_mcp = active_mcp.unwrap();
// MCP is active - transform tools and set up interception
handle_streaming_with_tool_interception(
client,
response_storage,
conversation_storage,
conversation_item_storage,
url,
headers,
payload,
original_body,
original_previous_response_id,
active_mcp,
)
.await
handle_streaming_with_tool_interception(&client, headers.as_ref(), req, &active_mcp).await
}