[model-gateway] introduce provider in openai router (#14394)

This commit is contained in:
Simo Lin
2025-12-03 22:47:20 -08:00
committed by GitHub
parent 654a78f9c8
commit fa78c44a3d
3 changed files with 406 additions and 81 deletions

View File

@@ -9,10 +9,12 @@
pub mod conversations;
pub mod mcp;
pub mod provider;
mod responses;
mod router;
mod streaming;
mod utils;
// Re-export the main router type for external use
// Re-export the main types for external use
pub use provider::{Provider, ProviderError, ProviderRegistry};
pub use router::OpenAIRouter;

View File

@@ -0,0 +1,356 @@
//! Provider abstractions for vendor-specific API transformations.
use std::{collections::HashMap, sync::Arc};
use reqwest::RequestBuilder;
use serde_json::Value;
use thiserror::Error;
use crate::core::{model_type::Endpoint, ProviderType};
const SGLANG_FIELDS: &[&str] = &[
"request_id",
"priority",
"top_k",
"min_p",
"min_tokens",
"regex",
"ebnf",
"json_schema",
"stop_token_ids",
"no_stop_trim",
"ignore_eos",
"continue_final_message",
"skip_special_tokens",
"lora_path",
"session_params",
"separate_reasoning",
"stream_reasoning",
"chat_template",
"chat_template_kwargs",
"return_hidden_states",
"repetition_penalty",
"sampling_seed",
"backend_url",
];
fn strip_sglang_fields(payload: &mut Value) {
if let Some(obj) = payload.as_object_mut() {
for field in SGLANG_FIELDS {
obj.remove(*field);
}
}
}
#[derive(Error, Debug)]
pub enum ProviderError {
#[error("Unsupported endpoint: {0:?}")]
UnsupportedEndpoint(Endpoint),
#[error("Transform error: {0}")]
TransformError(String),
}
/// Default `transform_request` strips SGLang fields.
pub trait Provider: Send + Sync {
fn provider_type(&self) -> ProviderType;
fn transform_request(
&self,
payload: &mut Value,
_endpoint: Endpoint,
) -> Result<(), ProviderError> {
strip_sglang_fields(payload);
Ok(())
}
fn transform_response(
&self,
_response: &mut Value,
_endpoint: Endpoint,
) -> Result<(), ProviderError> {
Ok(())
}
fn apply_headers(&self, builder: RequestBuilder) -> RequestBuilder {
builder
}
}
pub struct SGLangProvider;
impl Provider for SGLangProvider {
fn provider_type(&self) -> ProviderType {
ProviderType::OpenAI
}
fn transform_request(
&self,
_payload: &mut Value,
_endpoint: Endpoint,
) -> Result<(), ProviderError> {
Ok(())
}
}
pub struct OpenAIProvider;
impl Provider for OpenAIProvider {
fn provider_type(&self) -> ProviderType {
ProviderType::OpenAI
}
}
pub struct AnthropicProvider;
impl Provider for AnthropicProvider {
fn provider_type(&self) -> ProviderType {
ProviderType::Anthropic
}
}
pub struct XAIProvider;
impl Provider for XAIProvider {
fn provider_type(&self) -> ProviderType {
ProviderType::XAI
}
fn transform_request(
&self,
payload: &mut Value,
endpoint: Endpoint,
) -> Result<(), ProviderError> {
strip_sglang_fields(payload);
if endpoint == Endpoint::Responses {
if let Some(obj) = payload.as_object_mut() {
Self::transform_responses_input(obj);
}
}
Ok(())
}
}
impl XAIProvider {
fn transform_responses_input(obj: &mut serde_json::Map<String, Value>) {
let Some(input_arr) = obj.get_mut("input").and_then(Value::as_array_mut) else {
return;
};
for item in input_arr.iter_mut().filter_map(Value::as_object_mut) {
item.remove("id");
item.remove("status");
let Some(content_arr) = item.get_mut("content").and_then(Value::as_array_mut) else {
continue;
};
for content in content_arr.iter_mut().filter_map(Value::as_object_mut) {
if content.get("type").and_then(Value::as_str) == Some("output_text") {
content.insert("type".to_string(), Value::String("input_text".to_string()));
}
}
}
}
}
pub struct GeminiProvider;
impl Provider for GeminiProvider {
fn provider_type(&self) -> ProviderType {
ProviderType::Gemini
}
fn transform_request(
&self,
payload: &mut Value,
endpoint: Endpoint,
) -> Result<(), ProviderError> {
strip_sglang_fields(payload);
if endpoint == Endpoint::Chat {
if let Some(obj) = payload.as_object_mut() {
if obj.get("logprobs").and_then(|v| v.as_bool()) == Some(false) {
obj.remove("logprobs");
}
}
}
Ok(())
}
}
pub struct ProviderRegistry {
providers: HashMap<ProviderType, Arc<dyn Provider>>,
default_provider: Arc<dyn Provider>,
}
impl Default for ProviderRegistry {
fn default() -> Self {
Self::new()
}
}
impl ProviderRegistry {
pub fn new() -> Self {
let mut providers = HashMap::new();
providers.insert(
ProviderType::OpenAI,
Arc::new(OpenAIProvider) as Arc<dyn Provider>,
);
providers.insert(
ProviderType::XAI,
Arc::new(XAIProvider) as Arc<dyn Provider>,
);
providers.insert(
ProviderType::Gemini,
Arc::new(GeminiProvider) as Arc<dyn Provider>,
);
providers.insert(
ProviderType::Anthropic,
Arc::new(AnthropicProvider) as Arc<dyn Provider>,
);
Self {
providers,
default_provider: Arc::new(SGLangProvider),
}
}
pub fn get(&self, provider_type: &ProviderType) -> &dyn Provider {
self.providers
.get(provider_type)
.map(|p| p.as_ref())
.unwrap_or(self.default_provider.as_ref())
}
pub fn get_for_model(&self, model_name: &str) -> &dyn Provider {
match ProviderType::from_model_name(model_name) {
Some(pt) => self.get(&pt),
None => self.default_provider.as_ref(),
}
}
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
);
}
}

View File

@@ -14,7 +14,6 @@ use axum::{
Json,
};
use futures_util::{future::join_all, StreamExt};
use once_cell::sync::Lazy;
use serde_json::{json, to_value, Value};
use tokio::sync::mpsc;
use tokio_stream::wrappers::UnboundedReceiverStream;
@@ -31,13 +30,14 @@ use super::{
ensure_request_mcp_client, execute_tool_loop, prepare_mcp_payload_for_streaming,
McpLoopConfig,
},
provider::ProviderRegistry,
responses::{mask_tools_as_mcp, patch_streaming_response_json},
streaming::handle_streaming_response,
utils::{apply_provider_headers, extract_auth_header},
};
use crate::{
app_context::AppContext,
core::{ModelCard, RuntimeType, Worker, WorkerRegistry},
core::{model_type::Endpoint, ModelCard, ProviderType, RuntimeType, Worker, WorkerRegistry},
data_connector::{
ConversationId, ConversationItemStorage, ConversationStorage, ListParams, ResponseId,
ResponseStorage, SortOrder,
@@ -61,32 +61,6 @@ use crate::{
// OpenAIRouter Struct
// ============================================================================
/// Fields specific to SGLang that should be stripped when forwarding to OpenAI-compatible endpoints
static SGLANG_FIELDS: Lazy<HashSet<&'static str>> = Lazy::new(|| {
HashSet::from([
"request_id",
"priority",
"top_k",
"min_p",
"min_tokens",
"regex",
"ebnf",
"stop_token_ids",
"no_stop_trim",
"ignore_eos",
"continue_final_message",
"skip_special_tokens",
"lora_path",
"session_params",
"separate_reasoning",
"stream_reasoning",
"chat_template_kwargs",
"return_hidden_states",
"repetition_penalty",
"sampling_seed",
])
});
/// Router for OpenAI backend
///
/// This router manages connections to OpenAI-compatible API endpoints (OpenAI, xAI, etc.)
@@ -97,6 +71,8 @@ pub struct OpenAIRouter {
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
@@ -146,6 +122,7 @@ impl OpenAIRouter {
Ok(Self {
client,
worker_registry,
provider_registry: ProviderRegistry::new(),
healthy: AtomicBool::new(true),
response_storage: ctx.response_storage.clone(),
conversation_storage: ctx.conversation_storage.clone(),
@@ -154,6 +131,31 @@ impl OpenAIRouter {
})
}
/// 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,
worker: &dyn Worker,
model_id: Option<&str>,
) -> &'a dyn super::provider::Provider {
// Try worker's provider for the model first
if let Some(model) = model_id {
if let Some(pt) = worker.provider_for_model(model) {
return self.provider_registry.get(pt);
}
// Fall back to model name inference
if let Some(pt) = ProviderType::from_model_name(model) {
return self.provider_registry.get(&pt);
}
}
// Default to SGLang passthrough
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.
@@ -623,7 +625,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
&self,
headers: Option<&HeaderMap>,
body: &ChatCompletionRequest,
_model_id: Option<&str>,
model_id: Option<&str>,
) -> Response {
// Extract auth header for passthrough mode
let auth_header = extract_auth_header(headers, &None);
@@ -648,13 +650,14 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.into_response();
}
};
if let Some(obj) = payload.as_object_mut() {
// Always remove SGLang-specific fields (unsupported by OpenAI)
obj.retain(|k, _| !SGLANG_FIELDS.contains(&k.as_str()));
// Remove logprobs if false (Gemini don't accept it)
if obj.get("logprobs").and_then(|v| v.as_bool()) == Some(false) {
obj.remove("logprobs");
}
// Apply provider-specific transformations
let provider = self.get_provider_for_worker(worker.as_ref(), model_id);
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Chat) {
return (
StatusCode::BAD_REQUEST,
format!("Provider transform error: {}", e),
)
.into_response();
}
let url = format!("{}/v1/chat/completions", worker.url());
@@ -1019,50 +1022,14 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
};
// Remove SGLang-specific fields only
if let Some(obj) = payload.as_object_mut() {
// Remove SGLang-specific fields (not part of OpenAI API)
obj.retain(|k, _| !SGLANG_FIELDS.contains(&k.as_str()));
// XAI (Grok models) requires special handling of input items
// Check if model is a Grok model
let is_grok_model = obj
.get("model")
.and_then(|v| v.as_str())
.map(|m| m.starts_with("grok"))
.unwrap_or(false);
if is_grok_model {
// XAI doesn't support the OPENAI item type input: https://platform.openai.com/docs/api-reference/responses/create#responses-create-input-input-item-list-item
// To Achieve XAI compatibility, strip extra fields from input messages (id, status)
// XAI doesn't support output_text as type for content with role of assistant
// so normalize content types: output_text -> input_text
if let Some(input_arr) = obj.get_mut("input").and_then(Value::as_array_mut) {
for item_obj in input_arr.iter_mut().filter_map(Value::as_object_mut) {
// Remove fields not universally supported
item_obj.remove("id");
item_obj.remove("status");
// Normalize content types to input_text (xAI compatibility)
if let Some(content_arr) =
item_obj.get_mut("content").and_then(Value::as_array_mut)
{
for content_obj in
content_arr.iter_mut().filter_map(Value::as_object_mut)
{
// Change output_text to input_text
if content_obj.get("type").and_then(Value::as_str)
== Some("output_text")
{
content_obj.insert(
"type".to_string(),
Value::String("input_text".to_string()),
);
}
}
}
}
}
}
// Apply provider-specific transformations (handles SGLang fields, XAI/Grok, etc.)
let provider = self.get_provider_for_worker(worker.as_ref(), model_id);
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Responses) {
return (
StatusCode::BAD_REQUEST,
format!("Provider transform error: {}", e),
)
.into_response();
}
// Delegate to streaming or non-streaming handler