[model-gateway] introduce request ctx for oai router (#14434)
Co-authored-by: key4ng <rukeyang@gmail.com>
This commit is contained in:
@@ -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
|
//! - Multi-turn tool execution loops
|
||||||
//! - SSE (Server-Sent Events) streaming
|
//! - SSE (Server-Sent Events) streaming
|
||||||
|
|
||||||
|
mod context;
|
||||||
pub mod conversations;
|
pub mod conversations;
|
||||||
pub mod mcp;
|
pub mod mcp;
|
||||||
pub mod provider;
|
pub mod provider;
|
||||||
|
|||||||
@@ -225,6 +225,13 @@ impl ProviderRegistry {
|
|||||||
.unwrap_or(self.default_provider.as_ref())
|
.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 {
|
pub fn get_for_model(&self, model_name: &str) -> &dyn Provider {
|
||||||
match ProviderType::from_model_name(model_name) {
|
match ProviderType::from_model_name(model_name) {
|
||||||
Some(pt) => self.get(&pt),
|
Some(pt) => self.get(&pt),
|
||||||
@@ -235,122 +242,8 @@ impl ProviderRegistry {
|
|||||||
pub fn default_provider(&self) -> &dyn Provider {
|
pub fn default_provider(&self) -> &dyn Provider {
|
||||||
self.default_provider.as_ref()
|
self.default_provider.as_ref()
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
pub fn default_provider_arc(&self) -> Arc<dyn Provider> {
|
||||||
mod tests {
|
Arc::clone(&self.default_provider)
|
||||||
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
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
//! OpenAI router - main coordinator that delegates to specialized modules
|
|
||||||
|
|
||||||
use std::{
|
use std::{
|
||||||
any::Any,
|
any::Any,
|
||||||
collections::HashSet,
|
collections::HashSet,
|
||||||
@@ -19,13 +17,16 @@ use tokio::sync::mpsc;
|
|||||||
use tokio_stream::wrappers::UnboundedReceiverStream;
|
use tokio_stream::wrappers::UnboundedReceiverStream;
|
||||||
use tracing::warn;
|
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::{
|
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::{
|
mcp::{
|
||||||
ensure_request_mcp_client, execute_tool_loop, prepare_mcp_payload_for_streaming,
|
ensure_request_mcp_client, execute_tool_loop, prepare_mcp_payload_for_streaming,
|
||||||
McpLoopConfig,
|
McpLoopConfig,
|
||||||
@@ -38,11 +39,7 @@ use super::{
|
|||||||
use crate::{
|
use crate::{
|
||||||
app_context::AppContext,
|
app_context::AppContext,
|
||||||
core::{model_type::Endpoint, ModelCard, ProviderType, RuntimeType, Worker, WorkerRegistry},
|
core::{model_type::Endpoint, ModelCard, ProviderType, RuntimeType, Worker, WorkerRegistry},
|
||||||
data_connector::{
|
data_connector::{ConversationId, ListParams, ResponseId, SortOrder},
|
||||||
ConversationId, ConversationItemStorage, ConversationStorage, ListParams, ResponseId,
|
|
||||||
ResponseStorage, SortOrder,
|
|
||||||
},
|
|
||||||
mcp::McpManager,
|
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::ChatCompletionRequest,
|
chat::ChatCompletionRequest,
|
||||||
classify::ClassifyRequest,
|
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 {
|
pub struct OpenAIRouter {
|
||||||
/// HTTP client for upstream OpenAI-compatible API
|
|
||||||
client: reqwest::Client,
|
|
||||||
/// Worker registry for model-based worker lookup
|
|
||||||
worker_registry: Arc<WorkerRegistry>,
|
worker_registry: Arc<WorkerRegistry>,
|
||||||
/// Provider registry for vendor-specific transformations
|
|
||||||
provider_registry: ProviderRegistry,
|
provider_registry: ProviderRegistry,
|
||||||
/// Health status
|
|
||||||
healthy: AtomicBool,
|
healthy: AtomicBool,
|
||||||
/// Response storage for managing conversation history
|
shared_components: Arc<SharedComponents>,
|
||||||
response_storage: Arc<dyn ResponseStorage>,
|
responses_components: Arc<ResponsesComponents>,
|
||||||
/// 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>,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl std::fmt::Debug for OpenAIRouter {
|
impl std::fmt::Debug for OpenAIRouter {
|
||||||
@@ -98,76 +75,70 @@ impl std::fmt::Debug for OpenAIRouter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl OpenAIRouter {
|
impl OpenAIRouter {
|
||||||
/// Maximum number of conversation items to attach as input when a conversation is provided
|
|
||||||
const MAX_CONVERSATION_HISTORY_ITEMS: usize = 100;
|
const MAX_CONVERSATION_HISTORY_ITEMS: usize = 100;
|
||||||
|
|
||||||
/// Create a new OpenAI router
|
fn shared_components(&self) -> Arc<SharedComponents> {
|
||||||
///
|
Arc::clone(&self.shared_components)
|
||||||
/// Workers are registered separately via the external worker registration workflow.
|
}
|
||||||
/// This router queries the WorkerRegistry to find workers that support requested models.
|
|
||||||
|
fn responses_components(&self) -> Arc<ResponsesComponents> {
|
||||||
|
Arc::clone(&self.responses_components)
|
||||||
|
}
|
||||||
|
|
||||||
pub async fn new(ctx: &Arc<AppContext>) -> Result<Self, String> {
|
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();
|
let worker_registry = ctx.worker_registry.clone();
|
||||||
|
|
||||||
// Get MCP manager from AppContext (must be initialized)
|
|
||||||
let mcp_manager = ctx
|
let mcp_manager = ctx
|
||||||
.mcp_manager
|
.mcp_manager
|
||||||
.get()
|
.get()
|
||||||
.ok_or_else(|| "MCP manager not initialized in AppContext".to_string())?
|
.ok_or_else(|| "MCP manager not initialized in AppContext".to_string())?
|
||||||
.clone();
|
.clone();
|
||||||
|
|
||||||
Ok(Self {
|
let shared_components = Arc::new(SharedComponents {
|
||||||
client,
|
client: ctx.client.clone(),
|
||||||
worker_registry,
|
});
|
||||||
provider_registry: ProviderRegistry::new(),
|
|
||||||
healthy: AtomicBool::new(true),
|
let responses_components = Arc::new(ResponsesComponents {
|
||||||
|
shared: SharedComponents {
|
||||||
|
client: ctx.client.clone(),
|
||||||
|
},
|
||||||
|
mcp_manager: mcp_manager.clone(),
|
||||||
response_storage: ctx.response_storage.clone(),
|
response_storage: ctx.response_storage.clone(),
|
||||||
conversation_storage: ctx.conversation_storage.clone(),
|
conversation_storage: ctx.conversation_storage.clone(),
|
||||||
conversation_item_storage: ctx.conversation_item_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.
|
fn get_provider_arc_for_worker(
|
||||||
///
|
&self,
|
||||||
/// 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,
|
|
||||||
worker: &dyn Worker,
|
worker: &dyn Worker,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> &'a dyn super::provider::Provider {
|
) -> Arc<dyn super::provider::Provider> {
|
||||||
// Try worker's provider for the model first
|
|
||||||
if let Some(model) = model_id {
|
if let Some(model) = model_id {
|
||||||
if let Some(pt) = worker.provider_for_model(model) {
|
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) {
|
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_arc()
|
||||||
self.provider_registry.default_provider()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 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(
|
async fn refresh_worker_models(
|
||||||
&self,
|
&self,
|
||||||
worker: &Arc<dyn Worker>,
|
worker: &Arc<dyn Worker>,
|
||||||
auth_header: Option<&HeaderValue>,
|
auth_header: Option<&HeaderValue>,
|
||||||
) -> bool {
|
) -> bool {
|
||||||
let url = format!("{}/v1/models", worker.url());
|
let url = format!("{}/v1/models", worker.url());
|
||||||
|
let mut backend_req = self.shared_components.client.get(&url);
|
||||||
// Build request to backend
|
|
||||||
let mut backend_req = self.client.get(&url);
|
|
||||||
if let Some(auth) = auth_header {
|
if let Some(auth) = auth_header {
|
||||||
backend_req = apply_provider_headers(backend_req, &url, Some(auth));
|
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>) {
|
async fn refresh_external_models(&self, auth_header: Option<&HeaderValue>) {
|
||||||
let external_workers = self.worker_registry.get_workers_filtered(
|
let external_workers = self.worker_registry.get_workers_filtered(
|
||||||
None,
|
None,
|
||||||
@@ -235,7 +205,6 @@ impl OpenAIRouter {
|
|||||||
external_workers.len()
|
external_workers.len()
|
||||||
);
|
);
|
||||||
|
|
||||||
// Refresh all workers in parallel
|
|
||||||
let futures: Vec<_> = external_workers
|
let futures: Vec<_> = external_workers
|
||||||
.iter()
|
.iter()
|
||||||
.map(|w| self.refresh_worker_models(w, auth_header))
|
.map(|w| self.refresh_worker_models(w, auth_header))
|
||||||
@@ -244,41 +213,19 @@ impl OpenAIRouter {
|
|||||||
join_all(futures).await;
|
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(
|
async fn select_worker_for_model(
|
||||||
&self,
|
&self,
|
||||||
model_id: &str,
|
model_id: &str,
|
||||||
auth_header: Option<&HeaderValue>,
|
auth_header: Option<&HeaderValue>,
|
||||||
) -> Result<Arc<dyn Worker>, Box<Response>> {
|
) -> 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 = || {
|
let find_candidates = || {
|
||||||
self.worker_registry
|
self.worker_registry
|
||||||
.get_workers_filtered(
|
.get_workers_filtered(None, None, None, Some(RuntimeType::External), true)
|
||||||
None, // Get all external workers, not just those in model index
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
Some(RuntimeType::External),
|
|
||||||
true, // healthy_only
|
|
||||||
)
|
|
||||||
.into_iter()
|
.into_iter()
|
||||||
.filter(|w| w.supports_model(model_id) && w.circuit_breaker().can_execute())
|
.filter(|w| w.supports_model(model_id) && w.circuit_breaker().can_execute())
|
||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
};
|
};
|
||||||
|
|
||||||
// First try: find workers that already support this model
|
|
||||||
let candidates = find_candidates();
|
let candidates = find_candidates();
|
||||||
if !candidates.is_empty() {
|
if !candidates.is_empty() {
|
||||||
return Ok(candidates
|
return Ok(candidates
|
||||||
@@ -287,14 +234,12 @@ impl OpenAIRouter {
|
|||||||
.expect("candidates is not empty"));
|
.expect("candidates is not empty"));
|
||||||
}
|
}
|
||||||
|
|
||||||
// No match found - refresh models on all external workers
|
|
||||||
tracing::debug!(
|
tracing::debug!(
|
||||||
"No worker found for model '{}', refreshing external worker models",
|
"No worker found for model '{}', refreshing external worker models",
|
||||||
model_id
|
model_id
|
||||||
);
|
);
|
||||||
self.refresh_external_models(auth_header).await;
|
self.refresh_external_models(auth_header).await;
|
||||||
|
|
||||||
// Second try: check if any worker now supports the model after refresh
|
|
||||||
let candidates = find_candidates();
|
let candidates = find_candidates();
|
||||||
if !candidates.is_empty() {
|
if !candidates.is_empty() {
|
||||||
return Ok(candidates
|
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, mut ctx: RequestContext) -> Response {
|
||||||
async fn handle_non_streaming_response(
|
let payload_state = ctx.take_payload().expect("Payload not prepared");
|
||||||
&self,
|
let mut payload = payload_state.json;
|
||||||
worker: &Arc<dyn Worker>,
|
let url = payload_state.url;
|
||||||
headers: Option<&HeaderMap>,
|
let previous_response_id = payload_state.previous_response_id;
|
||||||
mut payload: Value,
|
let original_body = ctx.responses_request();
|
||||||
original_body: &ResponsesRequest,
|
let worker = ctx.worker().expect("Worker not selected");
|
||||||
original_previous_response_id: Option<String>,
|
let mcp_manager = ctx.components.mcp_manager().expect("MCP manager required");
|
||||||
) -> Response {
|
|
||||||
let url = format!("{}/v1/responses", worker.url());
|
|
||||||
|
|
||||||
// Check if MCP is active for this request
|
|
||||||
// Ensure dynamic client is created if needed
|
|
||||||
if let Some(ref tools) = original_body.tools {
|
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 mcp_manager.list_tools().is_empty() {
|
||||||
let active_mcp = if self.mcp_manager.list_tools().is_empty() {
|
|
||||||
None
|
None
|
||||||
} else {
|
} else {
|
||||||
Some(&self.mcp_manager)
|
Some(mcp_manager)
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut response_json: Value;
|
let mut response_json: Value;
|
||||||
|
|
||||||
// If MCP is active, execute tool loop
|
|
||||||
if let Some(mcp) = active_mcp {
|
if let Some(mcp) = active_mcp {
|
||||||
let config = McpLoopConfig::default();
|
let config = McpLoopConfig::default();
|
||||||
|
|
||||||
// Transform MCP tools to function tools
|
|
||||||
prepare_mcp_payload_for_streaming(&mut payload, mcp);
|
prepare_mcp_payload_for_streaming(&mut payload, mcp);
|
||||||
|
|
||||||
match execute_tool_loop(
|
match execute_tool_loop(
|
||||||
&self.client,
|
ctx.components.client(),
|
||||||
&url,
|
&url,
|
||||||
headers,
|
ctx.headers(),
|
||||||
payload,
|
payload,
|
||||||
original_body,
|
original_body,
|
||||||
mcp,
|
mcp,
|
||||||
@@ -372,13 +309,8 @@ impl OpenAIRouter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// No MCP - simple request
|
let mut request_builder = ctx.components.client().post(&url).json(&payload);
|
||||||
|
let auth_header = extract_auth_header(ctx.headers(), worker.api_key());
|
||||||
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());
|
|
||||||
request_builder = apply_provider_headers(request_builder, &url, auth_header.as_ref());
|
request_builder = apply_provider_headers(request_builder, &url, auth_header.as_ref());
|
||||||
|
|
||||||
let response = match request_builder.send().await {
|
let response = match request_builder.send().await {
|
||||||
@@ -421,19 +353,26 @@ impl OpenAIRouter {
|
|||||||
worker.circuit_breaker().record_success();
|
worker.circuit_breaker().record_success();
|
||||||
}
|
}
|
||||||
|
|
||||||
// Patch response with metadata
|
|
||||||
mask_tools_as_mcp(&mut response_json, original_body);
|
mask_tools_as_mcp(&mut response_json, original_body);
|
||||||
patch_streaming_response_json(
|
patch_streaming_response_json(
|
||||||
&mut response_json,
|
&mut response_json,
|
||||||
original_body,
|
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(
|
if let Err(err) = persist_conversation_items(
|
||||||
self.conversation_storage.clone(),
|
ctx.components
|
||||||
self.conversation_item_storage.clone(),
|
.conversation_storage()
|
||||||
self.response_storage.clone(),
|
.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,
|
&response_json,
|
||||||
original_body,
|
original_body,
|
||||||
)
|
)
|
||||||
@@ -446,10 +385,6 @@ impl OpenAIRouter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// ============================================================================
|
|
||||||
// RouterTrait Implementation
|
|
||||||
// ============================================================================
|
|
||||||
|
|
||||||
#[async_trait::async_trait]
|
#[async_trait::async_trait]
|
||||||
impl crate::routers::RouterTrait for OpenAIRouter {
|
impl crate::routers::RouterTrait for OpenAIRouter {
|
||||||
fn as_any(&self) -> &dyn Any {
|
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 {
|
async fn health_generate(&self, _req: Request<Body>) -> Response {
|
||||||
// Check health of all external workers
|
|
||||||
let external_workers: Vec<_> = self
|
let external_workers: Vec<_> = self
|
||||||
.worker_registry
|
.worker_registry
|
||||||
.get_all()
|
.get_all()
|
||||||
@@ -530,7 +464,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn get_models(&self, req: Request<Body>) -> Response {
|
async fn get_models(&self, req: Request<Body>) -> Response {
|
||||||
// Return models from all registered external workers
|
|
||||||
let external_workers: Vec<_> = self
|
let external_workers: Vec<_> = self
|
||||||
.worker_registry
|
.worker_registry
|
||||||
.get_all()
|
.get_all()
|
||||||
@@ -546,11 +479,9 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
.into_response();
|
.into_response();
|
||||||
}
|
}
|
||||||
|
|
||||||
// Refresh models for all external workers using user's auth header
|
|
||||||
let auth_header = extract_auth_header(Some(req.headers()), &None);
|
let auth_header = extract_auth_header(Some(req.headers()), &None);
|
||||||
self.refresh_external_models(auth_header.as_ref()).await;
|
self.refresh_external_models(auth_header.as_ref()).await;
|
||||||
|
|
||||||
// Collect models from all workers
|
|
||||||
let mut all_models = Vec::new();
|
let mut all_models = Vec::new();
|
||||||
let mut seen_models = HashSet::new();
|
let mut seen_models = HashSet::new();
|
||||||
|
|
||||||
@@ -562,7 +493,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
.map(|p| format!("{:?}", p).to_lowercase())
|
.map(|p| format!("{:?}", p).to_lowercase())
|
||||||
.unwrap_or_else(|| "unknown".to_string());
|
.unwrap_or_else(|| "unknown".to_string());
|
||||||
|
|
||||||
// Add primary model ID
|
|
||||||
if seen_models.insert(model_card.id.clone()) {
|
if seen_models.insert(model_card.id.clone()) {
|
||||||
all_models.push(json!({
|
all_models.push(json!({
|
||||||
"id": &model_card.id,
|
"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 {
|
for alias in &model_card.aliases {
|
||||||
if seen_models.insert(alias.clone()) {
|
if seen_models.insert(alias.clone()) {
|
||||||
all_models.push(json!({
|
all_models.push(json!({
|
||||||
@@ -589,7 +518,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Return aggregated models
|
|
||||||
let response_json = json!({
|
let response_json = json!({
|
||||||
"object": "list",
|
"object": "list",
|
||||||
"data": all_models
|
"data": all_models
|
||||||
@@ -599,7 +527,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn get_model_info(&self, _req: Request<Body>) -> Response {
|
async fn get_model_info(&self, _req: Request<Body>) -> Response {
|
||||||
// Not directly supported without model param; return 501
|
|
||||||
(
|
(
|
||||||
StatusCode::NOT_IMPLEMENTED,
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
"get_model_info not implemented for OpenAI router",
|
"get_model_info not implemented for OpenAI router",
|
||||||
@@ -613,7 +540,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
_body: &GenerateRequest,
|
_body: &GenerateRequest,
|
||||||
_model_id: Option<&str>,
|
_model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
// Generate endpoint is SGLang-specific, not supported for OpenAI backend
|
|
||||||
(
|
(
|
||||||
StatusCode::NOT_IMPLEMENTED,
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
"Generate endpoint not supported for OpenAI backend",
|
"Generate endpoint not supported for OpenAI backend",
|
||||||
@@ -627,10 +553,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
body: &ChatCompletionRequest,
|
body: &ChatCompletionRequest,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
// Extract auth header for passthrough mode
|
|
||||||
let auth_header = extract_auth_header(headers, &None);
|
let auth_header = extract_auth_header(headers, &None);
|
||||||
|
|
||||||
// Select worker for model (discovery happens inside if needed)
|
|
||||||
let worker = match self
|
let worker = match self
|
||||||
.select_worker_for_model(body.model.as_str(), auth_header.as_ref())
|
.select_worker_for_model(body.model.as_str(), auth_header.as_ref())
|
||||||
.await
|
.await
|
||||||
@@ -639,7 +563,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
Err(response) => return *response,
|
Err(response) => return *response,
|
||||||
};
|
};
|
||||||
|
|
||||||
// Serialize request body, removing SGLang-only fields
|
|
||||||
let mut payload = match to_value(body) {
|
let mut payload = match to_value(body) {
|
||||||
Ok(v) => v,
|
Ok(v) => v,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -650,8 +573,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
.into_response();
|
.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) {
|
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Chat) {
|
||||||
return (
|
return (
|
||||||
StatusCode::BAD_REQUEST,
|
StatusCode::BAD_REQUEST,
|
||||||
@@ -660,16 +583,31 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
.into_response();
|
.into_response();
|
||||||
}
|
}
|
||||||
|
|
||||||
let url = format!("{}/v1/chat/completions", worker.url());
|
let mut ctx = RequestContext::for_chat(
|
||||||
let mut req = self.client.post(&url).json(&payload);
|
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.)
|
ctx.state.worker = Some(WorkerSelection {
|
||||||
// Passthrough mode: user's auth header takes priority, worker's key is fallback
|
worker: Arc::clone(&worker),
|
||||||
let auth_header = extract_auth_header(headers, worker.api_key());
|
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());
|
req = apply_provider_headers(req, &url, auth_header.as_ref());
|
||||||
|
|
||||||
// Accept SSE when stream=true
|
if ctx.is_streaming() {
|
||||||
if body.stream {
|
|
||||||
req = req.header("Accept", "text/event-stream");
|
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())
|
let status = StatusCode::from_u16(resp.status().as_u16())
|
||||||
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
|
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
|
||||||
|
|
||||||
if !body.stream {
|
if !ctx.is_streaming() {
|
||||||
// Capture Content-Type before consuming response body
|
|
||||||
let content_type = resp.headers().get(CONTENT_TYPE).cloned();
|
let content_type = resp.headers().get(CONTENT_TYPE).cloned();
|
||||||
match resp.bytes().await {
|
match resp.bytes().await {
|
||||||
Ok(body) => {
|
Ok(body) => {
|
||||||
@@ -711,7 +648,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// Stream SSE bytes to client
|
|
||||||
let stream = resp.bytes_stream();
|
let stream = resp.bytes_stream();
|
||||||
let (tx, rx) = mpsc::unbounded_channel();
|
let (tx, rx) = mpsc::unbounded_channel();
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
@@ -745,10 +681,9 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
_body: &CompletionRequest,
|
_body: &CompletionRequest,
|
||||||
_model_id: Option<&str>,
|
_model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
// Completion endpoint not implemented for OpenAI backend
|
|
||||||
(
|
(
|
||||||
StatusCode::NOT_IMPLEMENTED,
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
"Completion endpoint not implemented for OpenAI backend",
|
"Completion endpoint not implemented",
|
||||||
)
|
)
|
||||||
.into_response()
|
.into_response()
|
||||||
}
|
}
|
||||||
@@ -759,10 +694,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
body: &ResponsesRequest,
|
body: &ResponsesRequest,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
// Extract auth header for passthrough mode
|
|
||||||
let auth_header = extract_auth_header(headers, &None);
|
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 model = model_id.unwrap_or(body.model.as_str());
|
||||||
let worker = match self
|
let worker = match self
|
||||||
.select_worker_for_model(model, auth_header.as_ref())
|
.select_worker_for_model(model, auth_header.as_ref())
|
||||||
@@ -772,22 +705,19 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
Err(response) => return *response,
|
Err(response) => return *response,
|
||||||
};
|
};
|
||||||
|
|
||||||
// Clone the body for validation and logic, but we'll build payload differently
|
|
||||||
let mut request_body = body.clone();
|
let mut request_body = body.clone();
|
||||||
if let Some(model) = model_id {
|
if let Some(model) = model_id {
|
||||||
request_body.model = model.to_string();
|
request_body.model = model.to_string();
|
||||||
}
|
}
|
||||||
// Do not forward conversation field upstream; retain for local persistence only
|
|
||||||
request_body.conversation = None;
|
request_body.conversation = None;
|
||||||
|
|
||||||
// Store the original previous_response_id for the response
|
|
||||||
let original_previous_response_id = request_body.previous_response_id.clone();
|
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;
|
let mut conversation_items: Option<Vec<ResponseInputOutputItem>> = None;
|
||||||
if let Some(prev_id_str) = request_body.previous_response_id.clone() {
|
if let Some(prev_id_str) = request_body.previous_response_id.clone() {
|
||||||
let prev_id = ResponseId::from(prev_id_str.as_str());
|
let prev_id = ResponseId::from(prev_id_str.as_str());
|
||||||
match self
|
match self
|
||||||
|
.responses_components
|
||||||
.response_storage
|
.response_storage
|
||||||
.get_response_chain(&prev_id, None)
|
.get_response_chain(&prev_id, None)
|
||||||
.await
|
.await
|
||||||
@@ -795,7 +725,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
Ok(chain) => {
|
Ok(chain) => {
|
||||||
let mut items = Vec::new();
|
let mut items = Vec::new();
|
||||||
for stored in chain.responses.iter() {
|
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() {
|
if let Some(input_arr) = stored.input.as_array() {
|
||||||
for item in input_arr {
|
for item in input_arr {
|
||||||
match serde_json::from_value::<ResponseInputOutputItem>(
|
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() {
|
if let Some(output_arr) = stored.output.as_array() {
|
||||||
for item in output_arr {
|
for item in output_arr {
|
||||||
match serde_json::from_value::<ResponseInputOutputItem>(
|
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() {
|
if let Some(conv_id_str) = body.conversation.clone() {
|
||||||
let conv_id = ConversationId::from(conv_id_str.as_str());
|
let conv_id = ConversationId::from(conv_id_str.as_str());
|
||||||
|
|
||||||
// Verify conversation exists
|
if let Ok(None) = self
|
||||||
if let Ok(None) = self.conversation_storage.get_conversation(&conv_id).await {
|
.responses_components
|
||||||
|
.conversation_storage
|
||||||
|
.get_conversation(&conv_id)
|
||||||
|
.await
|
||||||
|
{
|
||||||
return (
|
return (
|
||||||
StatusCode::NOT_FOUND,
|
StatusCode::NOT_FOUND,
|
||||||
Json(json!({"error": "Conversation not found"})),
|
Json(json!({"error": "Conversation not found"})),
|
||||||
@@ -855,7 +786,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
.into_response();
|
.into_response();
|
||||||
}
|
}
|
||||||
|
|
||||||
// Load conversation history (ascending order for chronological context)
|
|
||||||
let params = ListParams {
|
let params = ListParams {
|
||||||
limit: Self::MAX_CONVERSATION_HISTORY_ITEMS,
|
limit: Self::MAX_CONVERSATION_HISTORY_ITEMS,
|
||||||
order: SortOrder::Asc,
|
order: SortOrder::Asc,
|
||||||
@@ -863,6 +793,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
};
|
};
|
||||||
|
|
||||||
match self
|
match self
|
||||||
|
.responses_components
|
||||||
.conversation_item_storage
|
.conversation_item_storage
|
||||||
.list_items(&conv_id, params)
|
.list_items(&conv_id, params)
|
||||||
.await
|
.await
|
||||||
@@ -870,8 +801,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
Ok(stored_items) => {
|
Ok(stored_items) => {
|
||||||
let mut items: Vec<ResponseInputOutputItem> = Vec::new();
|
let mut items: Vec<ResponseInputOutputItem> = Vec::new();
|
||||||
for item in stored_items.into_iter() {
|
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() {
|
match item.item_type.as_str() {
|
||||||
"message" => {
|
"message" => {
|
||||||
match serde_json::from_value::<Vec<ResponseContentPart>>(
|
match serde_json::from_value::<Vec<ResponseContentPart>>(
|
||||||
@@ -897,7 +826,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
"function_call" => {
|
"function_call" => {
|
||||||
// The entire function_call item is stored in content field
|
|
||||||
match serde_json::from_value::<ResponseInputOutputItem>(
|
match serde_json::from_value::<ResponseInputOutputItem>(
|
||||||
item.content.clone(),
|
item.content.clone(),
|
||||||
) {
|
) {
|
||||||
@@ -911,7 +839,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
"function_call_output" => {
|
"function_call_output" => {
|
||||||
// The entire function_call_output item is stored in content field
|
|
||||||
tracing::debug!(
|
tracing::debug!(
|
||||||
"Loading function_call_output from DB - content: {}",
|
"Loading function_call_output from DB - content: {}",
|
||||||
serde_json::to_string_pretty(&item.content)
|
serde_json::to_string_pretty(&item.content)
|
||||||
@@ -934,17 +861,13 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
"reasoning" => {
|
"reasoning" => {}
|
||||||
// Skip reasoning items - they're internal processing details
|
|
||||||
}
|
|
||||||
_ => {
|
_ => {
|
||||||
// Skip unknown item types
|
|
||||||
warn!("Unknown item type in conversation: {}", item.item_type);
|
warn!("Unknown item type in conversation: {}", item.item_type);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Append current request
|
|
||||||
match &request_body.input {
|
match &request_body.input {
|
||||||
ResponseInput::Text(text) => {
|
ResponseInput::Text(text) => {
|
||||||
items.push(ResponseInputOutputItem::Message {
|
items.push(ResponseInputOutputItem::Message {
|
||||||
@@ -957,7 +880,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
ResponseInput::Items(current_items) => {
|
ResponseInput::Items(current_items) => {
|
||||||
// Process all item types, converting SimpleInputMessage to Message
|
|
||||||
for item in current_items.iter() {
|
for item in current_items.iter() {
|
||||||
let normalized =
|
let normalized =
|
||||||
crate::protocols::responses::normalize_input_item(item);
|
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 {
|
if let Some(mut items) = conversation_items {
|
||||||
// Append current request
|
|
||||||
match &request_body.input {
|
match &request_body.input {
|
||||||
ResponseInput::Text(text) => {
|
ResponseInput::Text(text) => {
|
||||||
items.push(ResponseInputOutputItem::Message {
|
items.push(ResponseInputOutputItem::Message {
|
||||||
@@ -992,7 +912,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
ResponseInput::Items(current_items) => {
|
ResponseInput::Items(current_items) => {
|
||||||
// Process all item types, converting SimpleInputMessage to Message
|
|
||||||
for item in current_items.iter() {
|
for item in current_items.iter() {
|
||||||
let normalized = crate::protocols::responses::normalize_input_item(item);
|
let normalized = crate::protocols::responses::normalize_input_item(item);
|
||||||
items.push(normalized);
|
items.push(normalized);
|
||||||
@@ -1003,14 +922,11 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
request_body.input = ResponseInput::Items(items);
|
request_body.input = ResponseInput::Items(items);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Always set store=false for upstream (we store internally)
|
|
||||||
request_body.store = Some(false);
|
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 {
|
if let ResponseInput::Items(ref mut items) = request_body.input {
|
||||||
items.retain(|item| !matches!(item, ResponseInputOutputItem::Reasoning { .. }));
|
items.retain(|item| !matches!(item, ResponseInputOutputItem::Reasoning { .. }));
|
||||||
}
|
}
|
||||||
|
|
||||||
// Convert to JSON and strip SGLang-specific fields
|
|
||||||
let mut payload = match to_value(&request_body) {
|
let mut payload = match to_value(&request_body) {
|
||||||
Ok(v) => v,
|
Ok(v) => v,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -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_arc_for_worker(worker.as_ref(), model_id);
|
||||||
let provider = self.get_provider_for_worker(worker.as_ref(), model_id);
|
|
||||||
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Responses) {
|
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Responses) {
|
||||||
return (
|
return (
|
||||||
StatusCode::BAD_REQUEST,
|
StatusCode::BAD_REQUEST,
|
||||||
@@ -1032,32 +947,28 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
.into_response();
|
.into_response();
|
||||||
}
|
}
|
||||||
|
|
||||||
// Delegate to streaming or non-streaming handler
|
let mut ctx = RequestContext::for_responses(
|
||||||
let url = format!("{}/v1/responses", worker.url());
|
Arc::new(body.clone()),
|
||||||
if body.stream.unwrap_or(false) {
|
headers.cloned(),
|
||||||
handle_streaming_response(
|
model_id.map(String::from),
|
||||||
&self.client,
|
ComponentRefs::Responses(self.responses_components()),
|
||||||
worker.circuit_breaker(),
|
);
|
||||||
Some(&self.mcp_manager),
|
|
||||||
self.response_storage.clone(),
|
ctx.state.worker = Some(WorkerSelection {
|
||||||
self.conversation_storage.clone(),
|
worker: Arc::clone(&worker),
|
||||||
self.conversation_item_storage.clone(),
|
provider: Arc::clone(&provider),
|
||||||
url,
|
});
|
||||||
headers,
|
|
||||||
payload,
|
ctx.state.payload = Some(PayloadState {
|
||||||
body,
|
json: payload,
|
||||||
original_previous_response_id,
|
url: format!("{}/v1/responses", worker.url()),
|
||||||
)
|
previous_response_id: original_previous_response_id,
|
||||||
.await
|
});
|
||||||
|
|
||||||
|
if ctx.is_streaming() {
|
||||||
|
handle_streaming_response(ctx).await
|
||||||
} else {
|
} else {
|
||||||
self.handle_non_streaming_response(
|
self.handle_non_streaming_response(ctx).await
|
||||||
&worker,
|
|
||||||
headers,
|
|
||||||
payload,
|
|
||||||
body,
|
|
||||||
original_previous_response_id,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1068,7 +979,12 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
_params: &ResponsesGetParams,
|
_params: &ResponsesGetParams,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
let id = ResponseId::from(response_id);
|
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)) => {
|
Ok(Some(stored)) => {
|
||||||
let mut response_json = stored.raw_response;
|
let mut response_json = stored.raw_response;
|
||||||
if let Some(obj) = response_json.as_object_mut() {
|
if let Some(obj) = response_json.as_object_mut() {
|
||||||
@@ -1104,20 +1020,22 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
) -> Response {
|
) -> Response {
|
||||||
let resp_id = ResponseId::from(response_id);
|
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)) => {
|
Ok(Some(stored)) => {
|
||||||
// Extract items from input field (which is a JSON array)
|
|
||||||
let items = match &stored.input {
|
let items = match &stored.input {
|
||||||
Value::Array(arr) => arr.clone(),
|
Value::Array(arr) => arr.clone(),
|
||||||
_ => vec![],
|
_ => vec![],
|
||||||
};
|
};
|
||||||
|
|
||||||
// Generate IDs for items if they don't have them
|
|
||||||
let items_with_ids: Vec<Value> = items
|
let items_with_ids: Vec<Value> = items
|
||||||
.into_iter()
|
.into_iter()
|
||||||
.map(|mut item| {
|
.map(|mut item| {
|
||||||
if item.get("id").is_none() {
|
if item.get("id").is_none() {
|
||||||
// Generate ID if not present using centralized utility
|
|
||||||
if let Some(obj) = item.as_object_mut() {
|
if let Some(obj) = item.as_object_mut() {
|
||||||
obj.insert("id".to_string(), json!(generate_id("msg")));
|
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 {
|
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(
|
async fn get_conversation(
|
||||||
@@ -1202,7 +1124,11 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
_headers: Option<&HeaderMap>,
|
_headers: Option<&HeaderMap>,
|
||||||
conversation_id: &str,
|
conversation_id: &str,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
get_conversation(&self.conversation_storage, conversation_id).await
|
get_conversation(
|
||||||
|
&self.responses_components.conversation_storage,
|
||||||
|
conversation_id,
|
||||||
|
)
|
||||||
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn update_conversation(
|
async fn update_conversation(
|
||||||
@@ -1211,7 +1137,12 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
conversation_id: &str,
|
conversation_id: &str,
|
||||||
body: &Value,
|
body: &Value,
|
||||||
) -> Response {
|
) -> 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(
|
async fn delete_conversation(
|
||||||
@@ -1219,7 +1150,11 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
_headers: Option<&HeaderMap>,
|
_headers: Option<&HeaderMap>,
|
||||||
conversation_id: &str,
|
conversation_id: &str,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
delete_conversation(&self.conversation_storage, conversation_id).await
|
delete_conversation(
|
||||||
|
&self.responses_components.conversation_storage,
|
||||||
|
conversation_id,
|
||||||
|
)
|
||||||
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn list_conversation_items(
|
async fn list_conversation_items(
|
||||||
@@ -1242,8 +1177,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
list_conversation_items(
|
list_conversation_items(
|
||||||
&self.conversation_storage,
|
&self.responses_components.conversation_storage,
|
||||||
&self.conversation_item_storage,
|
&self.responses_components.conversation_item_storage,
|
||||||
conversation_id,
|
conversation_id,
|
||||||
query_params,
|
query_params,
|
||||||
)
|
)
|
||||||
@@ -1257,8 +1192,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
body: &Value,
|
body: &Value,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
create_conversation_items(
|
create_conversation_items(
|
||||||
&self.conversation_storage,
|
&self.responses_components.conversation_storage,
|
||||||
&self.conversation_item_storage,
|
&self.responses_components.conversation_item_storage,
|
||||||
conversation_id,
|
conversation_id,
|
||||||
body.clone(),
|
body.clone(),
|
||||||
)
|
)
|
||||||
@@ -1273,8 +1208,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
include: Option<Vec<String>>,
|
include: Option<Vec<String>>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
get_conversation_item(
|
get_conversation_item(
|
||||||
&self.conversation_storage,
|
&self.responses_components.conversation_storage,
|
||||||
&self.conversation_item_storage,
|
&self.responses_components.conversation_item_storage,
|
||||||
conversation_id,
|
conversation_id,
|
||||||
item_id,
|
item_id,
|
||||||
include,
|
include,
|
||||||
@@ -1289,8 +1224,8 @@ impl crate::routers::RouterTrait for OpenAIRouter {
|
|||||||
item_id: &str,
|
item_id: &str,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
delete_conversation_item(
|
delete_conversation_item(
|
||||||
&self.conversation_storage,
|
&self.responses_components.conversation_storage,
|
||||||
&self.conversation_item_storage,
|
&self.responses_components.conversation_item_storage,
|
||||||
conversation_id,
|
conversation_id,
|
||||||
item_id,
|
item_id,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -22,8 +22,9 @@ use tokio_stream::wrappers::UnboundedReceiverStream;
|
|||||||
use tracing::warn;
|
use tracing::warn;
|
||||||
|
|
||||||
// Import from sibling modules
|
// Import from sibling modules
|
||||||
use super::conversations::persist_conversation_items;
|
use super::context::{RequestContext, StreamingEventContext, StreamingRequest};
|
||||||
use super::{
|
use super::{
|
||||||
|
conversations::persist_conversation_items,
|
||||||
mcp::{
|
mcp::{
|
||||||
build_resume_payload, ensure_request_mcp_client, execute_streaming_tool_calls,
|
build_resume_payload, ensure_request_mcp_client, execute_streaming_tool_calls,
|
||||||
inject_mcp_metadata_streaming, prepare_mcp_payload_for_streaming,
|
inject_mcp_metadata_streaming, prepare_mcp_payload_for_streaming,
|
||||||
@@ -33,7 +34,6 @@ use super::{
|
|||||||
utils::{event_types, FunctionCallInProgress, OutputIndexMapper, StreamAction},
|
utils::{event_types, FunctionCallInProgress, OutputIndexMapper, StreamAction},
|
||||||
};
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
data_connector::{ConversationItemStorage, ConversationStorage, ResponseStorage},
|
|
||||||
protocols::responses::{ResponseToolType, ResponsesRequest},
|
protocols::responses::{ResponseToolType, ResponsesRequest},
|
||||||
routers::header_utils::{apply_request_headers, preserve_response_headers},
|
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
|
/// Returns true if any changes were made
|
||||||
pub(super) fn apply_event_transformations_inplace(
|
pub(super) fn apply_event_transformations_inplace(
|
||||||
parsed_data: &mut Value,
|
parsed_data: &mut Value,
|
||||||
server_label: &str,
|
ctx: &StreamingEventContext<'_>,
|
||||||
original_request: &ResponsesRequest,
|
|
||||||
previous_response_id: Option<&str>,
|
|
||||||
) -> bool {
|
) -> bool {
|
||||||
let mut changed = false;
|
let mut changed = false;
|
||||||
|
|
||||||
@@ -575,13 +573,13 @@ pub(super) fn apply_event_transformations_inplace(
|
|||||||
.get_mut("response")
|
.get_mut("response")
|
||||||
.and_then(|v| v.as_object_mut())
|
.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) {
|
if response_obj.get("store") != Some(&desired_store) {
|
||||||
response_obj.insert("store".to_string(), desired_store);
|
response_obj.insert("store".to_string(), desired_store);
|
||||||
changed = true;
|
changed = true;
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(prev_id) = previous_response_id {
|
if let Some(prev_id) = ctx.previous_response_id {
|
||||||
let needs_previous = response_obj
|
let needs_previous = response_obj
|
||||||
.get("previous_response_id")
|
.get("previous_response_id")
|
||||||
.map(|v| v.is_null() || v.as_str().map(|s| s.is_empty()).unwrap_or(false))
|
.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)
|
// Mask tools from function to MCP format (optimized without cloning)
|
||||||
if response_obj.get("tools").is_some() {
|
if response_obj.get("tools").is_some() {
|
||||||
let requested_mcp = original_request
|
let requested_mcp = ctx
|
||||||
|
.original_request
|
||||||
.tools
|
.tools
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.map(|tools| {
|
.map(|tools| {
|
||||||
@@ -609,7 +608,7 @@ pub(super) fn apply_event_transformations_inplace(
|
|||||||
.unwrap_or(false);
|
.unwrap_or(false);
|
||||||
|
|
||||||
if requested_mcp {
|
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.insert("tools".to_string(), mcp_tools);
|
||||||
response_obj
|
response_obj
|
||||||
.entry("tool_choice".to_string())
|
.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 == event_types::ITEM_TYPE_FUNCTION_TOOL_CALL
|
||||||
{
|
{
|
||||||
item["type"] = json!(event_types::ITEM_TYPE_MCP_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_*
|
// Transform ID from fc_* to mcp_*
|
||||||
if let Some(id) = item.get("id").and_then(|v| v.as_str()) {
|
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
|
/// Forward and transform a streaming event to the client
|
||||||
/// Returns false if client disconnected
|
/// Returns false if client disconnected
|
||||||
#[allow(clippy::too_many_arguments)]
|
|
||||||
pub(super) fn forward_streaming_event(
|
pub(super) fn forward_streaming_event(
|
||||||
raw_block: &str,
|
raw_block: &str,
|
||||||
event_name: Option<&str>,
|
event_name: Option<&str>,
|
||||||
data: &str,
|
data: &str,
|
||||||
handler: &mut StreamingToolHandler,
|
handler: &mut StreamingToolHandler,
|
||||||
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
server_label: &str,
|
ctx: &StreamingEventContext<'_>,
|
||||||
original_request: &ResponsesRequest,
|
|
||||||
previous_response_id: Option<&str>,
|
|
||||||
sequence_number: &mut u64,
|
sequence_number: &mut u64,
|
||||||
) -> bool {
|
) -> bool {
|
||||||
// Skip individual function_call_arguments.delta events - we'll send them as one
|
// 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 all transformations in-place (single parse/serialize!)
|
||||||
apply_event_transformations_inplace(
|
apply_event_transformations_inplace(&mut parsed_data, ctx);
|
||||||
&mut parsed_data,
|
|
||||||
server_label,
|
|
||||||
original_request,
|
|
||||||
previous_response_id,
|
|
||||||
);
|
|
||||||
|
|
||||||
if let Some(response_obj) = parsed_data
|
if let Some(response_obj) = parsed_data
|
||||||
.get_mut("response")
|
.get_mut("response")
|
||||||
@@ -899,16 +890,13 @@ pub(super) fn forward_streaming_event(
|
|||||||
|
|
||||||
/// Send final response.completed event to client
|
/// Send final response.completed event to client
|
||||||
/// Returns false if client disconnected
|
/// Returns false if client disconnected
|
||||||
#[allow(clippy::too_many_arguments)]
|
|
||||||
pub(super) fn send_final_response_event(
|
pub(super) fn send_final_response_event(
|
||||||
handler: &StreamingToolHandler,
|
handler: &StreamingToolHandler,
|
||||||
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
sequence_number: &mut u64,
|
sequence_number: &mut u64,
|
||||||
state: &ToolLoopState,
|
state: &ToolLoopState,
|
||||||
active_mcp: Option<&Arc<crate::mcp::McpManager>>,
|
active_mcp: Option<&Arc<crate::mcp::McpManager>>,
|
||||||
original_request: &ResponsesRequest,
|
ctx: &StreamingEventContext<'_>,
|
||||||
previous_response_id: Option<&str>,
|
|
||||||
server_label: &str,
|
|
||||||
) -> bool {
|
) -> bool {
|
||||||
let mut final_response = match handler.snapshot_final_response() {
|
let mut final_response = match handler.snapshot_final_response() {
|
||||||
Some(resp) => resp,
|
Some(resp) => resp,
|
||||||
@@ -925,11 +913,15 @@ pub(super) fn send_final_response_event(
|
|||||||
}
|
}
|
||||||
|
|
||||||
if let Some(mcp) = active_mcp {
|
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);
|
mask_tools_as_mcp(&mut final_response, ctx.original_request);
|
||||||
patch_streaming_response_json(&mut final_response, original_request, previous_response_id);
|
patch_streaming_response_json(
|
||||||
|
&mut final_response,
|
||||||
|
ctx.original_request,
|
||||||
|
ctx.previous_response_id,
|
||||||
|
);
|
||||||
|
|
||||||
if let Some(obj) = final_response.as_object_mut() {
|
if let Some(obj) = final_response.as_object_mut() {
|
||||||
obj.insert("status".to_string(), Value::String("completed".to_string()));
|
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
|
/// Simple pass-through streaming without MCP interception
|
||||||
#[allow(clippy::too_many_arguments)]
|
|
||||||
pub(super) async fn handle_simple_streaming_passthrough(
|
pub(super) async fn handle_simple_streaming_passthrough(
|
||||||
client: &reqwest::Client,
|
client: &reqwest::Client,
|
||||||
circuit_breaker: &crate::core::CircuitBreaker,
|
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>,
|
headers: Option<&HeaderMap>,
|
||||||
payload: Value,
|
req: StreamingRequest,
|
||||||
original_body: &ResponsesRequest,
|
|
||||||
original_previous_response_id: Option<String>,
|
|
||||||
) -> Response {
|
) -> 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 {
|
if let Some(headers) = headers {
|
||||||
request_builder = apply_request_headers(headers, request_builder, true);
|
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 (tx, rx) = mpsc::unbounded_channel::<Result<Bytes, io::Error>>();
|
||||||
|
|
||||||
let should_store = original_body.store.unwrap_or(false);
|
let should_store = req.original_body.store.unwrap_or(false);
|
||||||
let original_request = original_body.clone();
|
let original_request = req.original_body;
|
||||||
let persist_needed = original_request.conversation.is_some();
|
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 {
|
tokio::spawn(async move {
|
||||||
let mut accumulator = StreamingResponseAccumulator::new();
|
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)
|
// Always persist conversation items and response (even without conversation)
|
||||||
if let Err(err) = persist_conversation_items(
|
if let Err(err) = persist_conversation_items(
|
||||||
conversation_storage.clone(),
|
storage.conversation.clone(),
|
||||||
conversation_item_storage.clone(),
|
storage.conversation_item.clone(),
|
||||||
response_storage.clone(),
|
storage.response.clone(),
|
||||||
&response_json,
|
&response_json,
|
||||||
&original_request,
|
&original_request,
|
||||||
)
|
)
|
||||||
@@ -1125,27 +1111,23 @@ pub(super) async fn handle_simple_streaming_passthrough(
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Handle streaming WITH MCP tool call interception and execution
|
/// Handle streaming WITH MCP tool call interception and execution
|
||||||
#[allow(clippy::too_many_arguments)]
|
|
||||||
pub(super) async fn handle_streaming_with_tool_interception(
|
pub(super) async fn handle_streaming_with_tool_interception(
|
||||||
client: &reqwest::Client,
|
client: &reqwest::Client,
|
||||||
response_storage: Arc<dyn ResponseStorage>,
|
|
||||||
conversation_storage: Arc<dyn ConversationStorage>,
|
|
||||||
conversation_item_storage: Arc<dyn ConversationItemStorage>,
|
|
||||||
url: String,
|
|
||||||
headers: Option<&HeaderMap>,
|
headers: Option<&HeaderMap>,
|
||||||
mut payload: Value,
|
req: StreamingRequest,
|
||||||
original_body: &ResponsesRequest,
|
|
||||||
original_previous_response_id: Option<String>,
|
|
||||||
active_mcp: &Arc<crate::mcp::McpManager>,
|
active_mcp: &Arc<crate::mcp::McpManager>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
// Transform MCP tools to function tools in payload
|
// Transform MCP tools to function tools in payload
|
||||||
|
let mut payload = req.payload;
|
||||||
prepare_mcp_payload_for_streaming(&mut payload, active_mcp);
|
prepare_mcp_payload_for_streaming(&mut payload, active_mcp);
|
||||||
|
|
||||||
let (tx, rx) = mpsc::unbounded_channel::<Result<Bytes, io::Error>>();
|
let (tx, rx) = mpsc::unbounded_channel::<Result<Bytes, io::Error>>();
|
||||||
let should_store = original_body.store.unwrap_or(false);
|
let should_store = req.original_body.store.unwrap_or(false);
|
||||||
let original_request = original_body.clone();
|
let original_request = req.original_body;
|
||||||
let persist_needed = original_request.conversation.is_some();
|
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 client_clone = client.clone();
|
||||||
let url_clone = url.clone();
|
let url_clone = url.clone();
|
||||||
@@ -1178,6 +1160,12 @@ pub(super) async fn handle_streaming_with_tool_interception(
|
|||||||
})
|
})
|
||||||
.unwrap_or("mcp");
|
.unwrap_or("mcp");
|
||||||
|
|
||||||
|
let streaming_ctx = StreamingEventContext {
|
||||||
|
server_label,
|
||||||
|
original_request: &original_request,
|
||||||
|
previous_response_id: previous_response_id.as_deref(),
|
||||||
|
};
|
||||||
|
|
||||||
loop {
|
loop {
|
||||||
// Make streaming request
|
// Make streaming request
|
||||||
let mut request_builder = client_clone.post(&url_clone).json(¤t_payload);
|
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(),
|
data.as_ref(),
|
||||||
&mut handler,
|
&mut handler,
|
||||||
&tx,
|
&tx,
|
||||||
server_label,
|
&streaming_ctx,
|
||||||
&original_request,
|
|
||||||
previous_response_id.as_deref(),
|
|
||||||
&mut sequence_number,
|
&mut sequence_number,
|
||||||
) {
|
) {
|
||||||
// Client disconnected
|
// Client disconnected
|
||||||
@@ -1319,9 +1305,7 @@ pub(super) async fn handle_streaming_with_tool_interception(
|
|||||||
data.as_ref(),
|
data.as_ref(),
|
||||||
&mut handler,
|
&mut handler,
|
||||||
&tx,
|
&tx,
|
||||||
server_label,
|
&streaming_ctx,
|
||||||
&original_request,
|
|
||||||
previous_response_id.as_deref(),
|
|
||||||
&mut sequence_number,
|
&mut sequence_number,
|
||||||
) {
|
) {
|
||||||
// Client disconnected
|
// Client disconnected
|
||||||
@@ -1358,9 +1342,7 @@ pub(super) async fn handle_streaming_with_tool_interception(
|
|||||||
&mut sequence_number,
|
&mut sequence_number,
|
||||||
&state,
|
&state,
|
||||||
Some(&active_mcp_clone),
|
Some(&active_mcp_clone),
|
||||||
&original_request,
|
&streaming_ctx,
|
||||||
previous_response_id.as_deref(),
|
|
||||||
server_label,
|
|
||||||
) {
|
) {
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -1393,9 +1375,9 @@ pub(super) async fn handle_streaming_with_tool_interception(
|
|||||||
|
|
||||||
// Always persist conversation items and response (even without conversation)
|
// Always persist conversation items and response (even without conversation)
|
||||||
if let Err(err) = persist_conversation_items(
|
if let Err(err) = persist_conversation_items(
|
||||||
conversation_storage.clone(),
|
storage.conversation.clone(),
|
||||||
conversation_item_storage.clone(),
|
storage.conversation_item.clone(),
|
||||||
response_storage.clone(),
|
storage.response.clone(),
|
||||||
&response_json,
|
&response_json,
|
||||||
&original_request,
|
&original_request,
|
||||||
)
|
)
|
||||||
@@ -1483,50 +1465,32 @@ pub(super) async fn handle_streaming_with_tool_interception(
|
|||||||
response
|
response
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Main entry point for handling streaming responses
|
pub(super) async fn handle_streaming_response(ctx: RequestContext) -> Response {
|
||||||
/// Delegates to simple passthrough or MCP tool interception based on configuration
|
let worker = ctx.worker().expect("Worker not selected").clone();
|
||||||
#[allow(clippy::too_many_arguments)]
|
let circuit_breaker = worker.circuit_breaker();
|
||||||
pub(super) async fn handle_streaming_response(
|
let headers = ctx.headers().cloned();
|
||||||
client: &reqwest::Client,
|
let original_body = ctx.responses_request();
|
||||||
circuit_breaker: &crate::core::CircuitBreaker,
|
let mcp_manager = ctx.components.mcp_manager().expect("MCP manager required");
|
||||||
mcp_manager: Option<&Arc<crate::mcp::McpManager>>,
|
|
||||||
response_storage: Arc<dyn ResponseStorage>,
|
if let Some(ref tools) = original_body.tools {
|
||||||
conversation_storage: Arc<dyn ConversationStorage>,
|
ensure_request_mcp_client(mcp_manager, tools.as_slice()).await;
|
||||||
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;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use the tool loop if the manager has any tools available (static or dynamic).
|
let active_mcp = if mcp_manager.list_tools().is_empty() {
|
||||||
let active_mcp = mcp_manager.and_then(|mgr| {
|
None
|
||||||
if mgr.list_tools().is_empty() {
|
} else {
|
||||||
None
|
Some(mcp_manager.clone())
|
||||||
} else {
|
};
|
||||||
Some(mgr)
|
|
||||||
}
|
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() {
|
if active_mcp.is_none() {
|
||||||
return handle_simple_streaming_passthrough(
|
return handle_simple_streaming_passthrough(
|
||||||
client,
|
&client,
|
||||||
circuit_breaker,
|
circuit_breaker,
|
||||||
response_storage,
|
headers.as_ref(),
|
||||||
conversation_storage,
|
req,
|
||||||
conversation_item_storage,
|
|
||||||
url,
|
|
||||||
headers,
|
|
||||||
payload,
|
|
||||||
original_body,
|
|
||||||
original_previous_response_id,
|
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
}
|
}
|
||||||
@@ -1534,17 +1498,5 @@ pub(super) async fn handle_streaming_response(
|
|||||||
let active_mcp = active_mcp.unwrap();
|
let active_mcp = active_mcp.unwrap();
|
||||||
|
|
||||||
// MCP is active - transform tools and set up interception
|
// MCP is active - transform tools and set up interception
|
||||||
handle_streaming_with_tool_interception(
|
handle_streaming_with_tool_interception(&client, headers.as_ref(), req, &active_mcp).await
|
||||||
client,
|
|
||||||
response_storage,
|
|
||||||
conversation_storage,
|
|
||||||
conversation_item_storage,
|
|
||||||
url,
|
|
||||||
headers,
|
|
||||||
payload,
|
|
||||||
original_body,
|
|
||||||
original_previous_response_id,
|
|
||||||
active_mcp,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user