diff --git a/sgl-router/src/routers/openai/context.rs b/sgl-router/src/routers/openai/context.rs new file mode 100644 index 000000000..462557fe6 --- /dev/null +++ b/sgl-router/src/routers/openai/context.rs @@ -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, + #[allow(dead_code)] + pub model_id: Option, +} + +pub enum RequestType { + Chat(Arc), + Responses(Arc), +} + +#[derive(Clone)] +pub struct SharedComponents { + pub client: reqwest::Client, +} + +pub struct ResponsesComponents { + pub shared: SharedComponents, + pub mcp_manager: Arc, + pub response_storage: Arc, + pub conversation_storage: Arc, + pub conversation_item_storage: Arc, +} + +pub enum ComponentRefs { + Shared(Arc), + Responses(Arc), +} + +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> { + match self { + ComponentRefs::Shared(_) => None, + ComponentRefs::Responses(r) => Some(&r.mcp_manager), + } + } + + pub fn response_storage(&self) -> Option<&Arc> { + match self { + ComponentRefs::Shared(_) => None, + ComponentRefs::Responses(r) => Some(&r.response_storage), + } + } + + pub fn conversation_storage(&self) -> Option<&Arc> { + match self { + ComponentRefs::Shared(_) => None, + ComponentRefs::Responses(r) => Some(&r.conversation_storage), + } + } + + pub fn conversation_item_storage(&self) -> Option<&Arc> { + match self { + ComponentRefs::Shared(_) => None, + ComponentRefs::Responses(r) => Some(&r.conversation_item_storage), + } + } +} + +#[derive(Default)] +pub struct ProcessingState { + pub worker: Option, + pub payload: Option, +} + +pub struct WorkerSelection { + pub worker: Arc, + #[allow(dead_code)] + pub provider: Arc, +} + +pub struct PayloadState { + pub json: Value, + pub url: String, + pub previous_response_id: Option, +} + +impl RequestContext { + pub fn for_responses( + request: Arc, + headers: Option, + model_id: Option, + components: ComponentRefs, + ) -> Self { + Self { + input: RequestInput { + request_type: RequestType::Responses(request), + headers, + model_id, + }, + components, + state: ProcessingState::default(), + } + } + + pub fn for_chat( + request: Arc, + headers: Option, + model_id: Option, + 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 { + 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> { + 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 { + self.state.payload.take() + } +} + +pub struct StorageHandles { + pub response: Arc, + pub conversation: Arc, + pub conversation_item: Arc, +} + +pub struct OwnedStreamingContext { + pub url: String, + pub payload: Value, + pub original_body: ResponsesRequest, + pub previous_response_id: Option, + 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; diff --git a/sgl-router/src/routers/openai/mod.rs b/sgl-router/src/routers/openai/mod.rs index 5e0eb5179..04358d73e 100644 --- a/sgl-router/src/routers/openai/mod.rs +++ b/sgl-router/src/routers/openai/mod.rs @@ -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; diff --git a/sgl-router/src/routers/openai/provider.rs b/sgl-router/src/routers/openai/provider.rs index 733611e14..545555579 100644 --- a/sgl-router/src/routers/openai/provider.rs +++ b/sgl-router/src/routers/openai/provider.rs @@ -225,6 +225,13 @@ impl ProviderRegistry { .unwrap_or(self.default_provider.as_ref()) } + pub fn get_arc(&self, provider_type: &ProviderType) -> Arc { + 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 { + Arc::clone(&self.default_provider) } } diff --git a/sgl-router/src/routers/openai/router.rs b/sgl-router/src/routers/openai/router.rs index 8a6fd8175..0c4c6aba9 100644 --- a/sgl-router/src/routers/openai/router.rs +++ b/sgl-router/src/routers/openai/router.rs @@ -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, - /// Provider registry for vendor-specific transformations provider_registry: ProviderRegistry, - /// Health status healthy: AtomicBool, - /// Response storage for managing conversation history - response_storage: Arc, - /// Conversation storage backend - conversation_storage: Arc, - /// Conversation item storage backend - conversation_item_storage: Arc, - /// MCP manager (handles both static and dynamic servers) - mcp_manager: Arc, + shared_components: Arc, + responses_components: Arc, } 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 { + Arc::clone(&self.shared_components) + } + + fn responses_components(&self) -> Arc { + Arc::clone(&self.responses_components) + } + pub async fn new(ctx: &Arc) -> Result { - // 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 { 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, 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, Box> { - // 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::>() }; - // 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, - headers: Option<&HeaderMap>, - mut payload: Value, - original_body: &ResponsesRequest, - original_previous_response_id: Option, - ) -> 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) -> 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) -> 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) -> 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> = 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::( @@ -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::( @@ -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 = 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::>( @@ -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::( 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 = 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>, ) -> 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, ) diff --git a/sgl-router/src/routers/openai/streaming.rs b/sgl-router/src/routers/openai/streaming.rs index bebda8eba..a1608931a 100644 --- a/sgl-router/src/routers/openai/streaming.rs +++ b/sgl-router/src/routers/openai/streaming.rs @@ -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 { /// 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>, - 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>, sequence_number: &mut u64, state: &ToolLoopState, active_mcp: Option<&Arc>, - 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, - conversation_storage: Arc, - conversation_item_storage: Arc, - url: String, headers: Option<&HeaderMap>, - payload: Value, - original_body: &ResponsesRequest, - original_previous_response_id: Option, + 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::>(); - 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, - conversation_storage: Arc, - conversation_item_storage: Arc, - url: String, headers: Option<&HeaderMap>, - mut payload: Value, - original_body: &ResponsesRequest, - original_previous_response_id: Option, + req: StreamingRequest, active_mcp: &Arc, ) -> 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::>(); - 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(¤t_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>, - response_storage: Arc, - conversation_storage: Arc, - conversation_item_storage: Arc, - url: String, - headers: Option<&HeaderMap>, - payload: Value, - original_body: &ResponsesRequest, - original_previous_response_id: Option, -) -> 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 }