[model-gateway] introduce request ctx for oai router (#14434)
Co-authored-by: key4ng <rukeyang@gmail.com>
This commit is contained in:
243
sgl-router/src/routers/openai/context.rs
Normal file
243
sgl-router/src/routers/openai/context.rs
Normal 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;
|
||||
@@ -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;
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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(¤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<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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user