[model-gateway] use worker crate in openai router (#14330)

This commit is contained in:
Simo Lin
2025-12-03 13:36:32 -08:00
committed by GitHub
parent 9d82340298
commit 388151053d
13 changed files with 612 additions and 384 deletions
+51 -1
View File
@@ -2,7 +2,7 @@ use std::{
fmt, fmt,
sync::{ sync::{
atomic::{AtomicBool, AtomicUsize, Ordering}, atomic::{AtomicBool, AtomicUsize, Ordering},
Arc, LazyLock, Arc, LazyLock, RwLock as StdRwLock,
}, },
time::{Duration, Instant}, time::{Duration, Instant},
}; };
@@ -260,6 +260,18 @@ pub trait Worker: Send + Sync + fmt::Debug {
&self.metadata().models &self.metadata().models
} }
/// Set models for this worker (for lazy discovery).
/// Default implementation does nothing - only BasicWorker supports this.
fn set_models(&self, _models: Vec<ModelCard>) {
// Default: no-op. BasicWorker overrides this.
}
/// Check if models have been discovered for this worker.
/// Returns true if models were set via set_models() or if metadata has models.
fn has_models_discovered(&self) -> bool {
!self.metadata().models.is_empty()
}
/// Get or create a gRPC client for this worker /// Get or create a gRPC client for this worker
/// Returns None for HTTP workers, Some(client) for gRPC workers /// Returns None for HTTP workers, Some(client) for gRPC workers
async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<GrpcClient>>>; async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<GrpcClient>>>;
@@ -485,6 +497,10 @@ pub struct BasicWorker {
pub circuit_breaker: CircuitBreaker, pub circuit_breaker: CircuitBreaker,
/// Lazily initialized gRPC client for gRPC workers /// Lazily initialized gRPC client for gRPC workers
pub grpc_client: Arc<RwLock<Option<Arc<GrpcClient>>>>, pub grpc_client: Arc<RwLock<Option<Arc<GrpcClient>>>>,
/// Runtime-mutable models override (for lazy discovery)
/// When set, overrides metadata.models for routing decisions.
/// Uses std::sync::RwLock for synchronous access in supports_model().
pub models_override: Arc<StdRwLock<Option<Vec<ModelCard>>>>,
} }
impl fmt::Debug for BasicWorker { impl fmt::Debug for BasicWorker {
@@ -622,6 +638,40 @@ impl Worker for BasicWorker {
&self.circuit_breaker &self.circuit_breaker
} }
fn supports_model(&self, model_id: &str) -> bool {
// Check models_override first (for lazy discovery)
if let Ok(guard) = self.models_override.read() {
if let Some(ref models) = *guard {
// Models were discovered - check if this model is supported
return models.iter().any(|m| m.matches(model_id));
}
}
// Fall back to metadata.models (empty = wildcard = supports nothing until discovery)
self.metadata.supports_model(model_id)
}
fn set_models(&self, models: Vec<ModelCard>) {
if let Ok(mut guard) = self.models_override.write() {
tracing::debug!(
"Setting {} models for worker {} via lazy discovery",
models.len(),
self.metadata.url
);
*guard = Some(models);
}
}
fn has_models_discovered(&self) -> bool {
// Check if models_override has been set
if let Ok(guard) = self.models_override.read() {
if guard.is_some() {
return true;
}
}
// Fall back to checking metadata.models
!self.metadata.models.is_empty()
}
async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<GrpcClient>>> { async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<GrpcClient>>> {
match self.metadata.connection_mode { match self.metadata.connection_mode {
ConnectionMode::Http => Ok(None), ConnectionMode::Http => Ok(None),
+2 -1
View File
@@ -128,7 +128,7 @@ impl BasicWorkerBuilder {
pub fn build(self) -> BasicWorker { pub fn build(self) -> BasicWorker {
use std::sync::{ use std::sync::{
atomic::{AtomicBool, AtomicUsize}, atomic::{AtomicBool, AtomicUsize},
Arc, Arc, RwLock as StdRwLock,
}; };
use tokio::sync::RwLock; use tokio::sync::RwLock;
@@ -187,6 +187,7 @@ impl BasicWorkerBuilder {
consecutive_successes: Arc::new(AtomicUsize::new(0)), consecutive_successes: Arc::new(AtomicUsize::new(0)),
circuit_breaker: CircuitBreaker::with_config(self.circuit_breaker_config), circuit_breaker: CircuitBreaker::with_config(self.circuit_breaker_config),
grpc_client, grpc_client,
models_override: Arc::new(StdRwLock::new(None)),
} }
} }
} }
+10 -1
View File
@@ -7,7 +7,7 @@ use std::sync::{Arc, RwLock};
use dashmap::DashMap; use dashmap::DashMap;
use uuid::Uuid; use uuid::Uuid;
use crate::core::{ConnectionMode, Worker, WorkerType}; use crate::core::{ConnectionMode, RuntimeType, Worker, WorkerType};
/// Unique identifier for a worker /// Unique identifier for a worker
#[derive(Debug, Clone, Hash, Eq, PartialEq)] #[derive(Debug, Clone, Hash, Eq, PartialEq)]
@@ -283,12 +283,14 @@ impl WorkerRegistry {
/// - model_id: Filter by specific model /// - model_id: Filter by specific model
/// - worker_type: Filter by worker type (Regular, Prefill, Decode) /// - worker_type: Filter by worker type (Regular, Prefill, Decode)
/// - connection_mode: Filter by connection mode (Http, Grpc) /// - connection_mode: Filter by connection mode (Http, Grpc)
/// - runtime_type: Filter by runtime type (Sglang, Vllm, External)
/// - healthy_only: Only return healthy workers /// - healthy_only: Only return healthy workers
pub fn get_workers_filtered( pub fn get_workers_filtered(
&self, &self,
model_id: Option<&str>, model_id: Option<&str>,
worker_type: Option<WorkerType>, worker_type: Option<WorkerType>,
connection_mode: Option<ConnectionMode>, connection_mode: Option<ConnectionMode>,
runtime_type: Option<RuntimeType>,
healthy_only: bool, healthy_only: bool,
) -> Vec<Arc<dyn Worker>> { ) -> Vec<Arc<dyn Worker>> {
// Start with the most efficient collection based on filters // Start with the most efficient collection based on filters
@@ -317,6 +319,13 @@ impl WorkerRegistry {
} }
} }
// Check runtime_type if specified
if let Some(ref rt) = runtime_type {
if w.metadata().runtime_type != *rt {
return false;
}
}
// Check health if required // Check health if required
if healthy_only && !w.is_healthy() { if healthy_only && !w.is_healthy() {
return false; return false;
@@ -222,6 +222,17 @@ impl StepExecutor for DiscoverModelsStep {
.get("worker_config") .get("worker_config")
.ok_or_else(|| WorkflowError::ContextValueNotFound("worker_config".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("worker_config".to_string()))?;
// If no API key is provided, skip model discovery and use wildcard mode.
if config.api_key.as_ref().is_none_or(|k| k.is_empty()) {
info!(
"No API key provided for {} - using wildcard mode (accepts any model). \
User's Authorization header will be forwarded to backend.",
config.url
);
context.set::<Vec<ModelCard>>("model_cards", vec![]);
return Ok(StepResult::Success);
}
debug!("Discovering models from external endpoint {}", config.url); debug!("Discovering models from external endpoint {}", config.url);
let model_cards = fetch_models(&config.url, config.api_key.as_deref()) let model_cards = fetch_models(&config.url, config.api_key.as_deref())
@@ -315,12 +326,6 @@ impl StepExecutor for CreateExternalWorkersStep {
.get("model_cards") .get("model_cards")
.ok_or_else(|| WorkflowError::ContextValueNotFound("model_cards".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("model_cards".to_string()))?;
debug!(
"Creating {} external workers for {}",
model_cards.len(),
config.url
);
// Build configs from router settings // Build configs from router settings
let circuit_breaker_config = { let circuit_breaker_config = {
let cfg = app_context.router_config.effective_circuit_breaker_config(); let cfg = app_context.router_config.effective_circuit_breaker_config();
@@ -355,8 +360,45 @@ impl StepExecutor for CreateExternalWorkersStep {
// Normalize URL (ensure https:// for external APIs) // Normalize URL (ensure https:// for external APIs)
let normalized_url = normalize_external_url(&config.url); let normalized_url = normalize_external_url(&config.url);
// Create a worker for each model
let mut workers = Vec::new(); let mut workers = Vec::new();
// Handle wildcard mode: create a single worker with empty models list
if model_cards.is_empty() {
debug!("Creating wildcard worker (no models) for {}", config.url);
let mut builder = BasicWorkerBuilder::new(normalized_url.clone())
.models(vec![]) // Empty models = accepts any model
.worker_type(WorkerType::Regular)
.connection_mode(ConnectionMode::Http)
.runtime_type(RuntimeType::External)
.circuit_breaker_config(circuit_breaker_config.clone())
.health_config(health_config.clone());
if let Some(ref api_key) = config.api_key {
builder = builder.api_key(api_key.clone());
}
if !labels.is_empty() {
builder = builder.labels(labels.clone());
}
let worker = Arc::new(builder.build()) as Arc<dyn Worker>;
worker.set_healthy(false);
info!(
"Created wildcard worker at {} (accepts any model, user auth forwarded)",
normalized_url
);
workers.push(worker);
} else {
debug!(
"Creating {} external workers for {}",
model_cards.len(),
config.url
);
// Create a worker for each model
for model_card in model_cards.iter() { for model_card in model_cards.iter() {
let mut builder = BasicWorkerBuilder::new(normalized_url.clone()) let mut builder = BasicWorkerBuilder::new(normalized_url.clone())
.model(model_card.clone()) .model(model_card.clone())
@@ -390,6 +432,7 @@ impl StepExecutor for CreateExternalWorkersStep {
workers.len(), workers.len(),
config.url config.url
); );
}
context.set("workers", workers); context.set("workers", workers);
context.set("labels", labels); context.set("labels", labels);
+7 -13
View File
@@ -56,9 +56,7 @@ impl RouterFactory {
) )
.await .await
} }
RoutingMode::OpenAI { worker_urls } => { RoutingMode::OpenAI { .. } => Self::create_openai_router(ctx).await,
Self::create_openai_router(worker_urls.clone(), ctx).await
}
}, },
} }
} }
@@ -119,16 +117,12 @@ impl RouterFactory {
} }
/// Create an OpenAI router /// Create an OpenAI router
async fn create_openai_router( ///
worker_urls: Vec<String>, /// Workers should be registered via the external worker registration workflow
ctx: &Arc<AppContext>, /// before using this router. The workflow discovers models from the provided
) -> Result<Box<dyn RouterTrait>, String> { /// endpoints and creates external workers in the registry.
if worker_urls.is_empty() { async fn create_openai_router(ctx: &Arc<AppContext>) -> Result<Box<dyn RouterTrait>, String> {
return Err("OpenAI mode requires at least one worker URL".to_string()); let router = OpenAIRouter::new(ctx).await?;
}
let router = OpenAIRouter::new(worker_urls, ctx).await?;
Ok(Box::new(router)) Ok(Box::new(router))
} }
} }
@@ -120,6 +120,7 @@ impl WorkerSelectionStage {
model_id, model_id,
Some(WorkerType::Regular), Some(WorkerType::Regular),
Some(ConnectionMode::Grpc { port: None }), Some(ConnectionMode::Grpc { port: None }),
None, // any runtime type
false, // get all workers, we'll filter by is_available() next false, // get all workers, we'll filter by is_available() next
); );
@@ -153,6 +154,7 @@ impl WorkerSelectionStage {
model_id, model_id,
None, None,
Some(ConnectionMode::Grpc { port: None }), // Match any gRPC worker Some(ConnectionMode::Grpc { port: None }), // Match any gRPC worker
None, // any runtime type
false, false,
); );
+2
View File
@@ -137,12 +137,14 @@ impl std::fmt::Debug for GrpcPDRouter {
bootstrap_port: None, bootstrap_port: None,
}), }),
Some(ConnectionMode::Grpc { port: None }), Some(ConnectionMode::Grpc { port: None }),
None,
false, false,
); );
let decode_workers = self.worker_registry.get_workers_filtered( let decode_workers = self.worker_registry.get_workers_filtered(
None, None,
Some(WorkerType::Decode), Some(WorkerType::Decode),
Some(ConnectionMode::Grpc { port: None }), Some(ConnectionMode::Grpc { port: None }),
None,
false, false,
); );
f.debug_struct("GrpcPDRouter") f.debug_struct("GrpcPDRouter")
+2
View File
@@ -53,6 +53,7 @@ impl Router {
None, // any model None, // any model
Some(WorkerType::Regular), Some(WorkerType::Regular),
Some(ConnectionMode::Http), Some(ConnectionMode::Http),
None, // any runtime type
false, // include all workers false, // include all workers
); );
@@ -139,6 +140,7 @@ impl Router {
effective_model_id, effective_model_id,
Some(WorkerType::Regular), Some(WorkerType::Regular),
Some(ConnectionMode::Http), Some(ConnectionMode::Http),
None, // any runtime type
false, // get all workers, we'll filter by is_available() next false, // get all workers, we'll filter by is_available() next
); );
+307 -232
View File
@@ -4,7 +4,6 @@ use std::{
any::Any, any::Any,
collections::HashSet, collections::HashSet,
sync::{atomic::AtomicBool, Arc}, sync::{atomic::AtomicBool, Arc},
time::{Duration, Instant},
}; };
use axum::{ use axum::{
@@ -14,8 +13,7 @@ use axum::{
response::{IntoResponse, Response}, response::{IntoResponse, Response},
Json, Json,
}; };
use dashmap::DashMap; use futures_util::{future::join_all, StreamExt};
use futures_util::StreamExt;
use once_cell::sync::Lazy; use once_cell::sync::Lazy;
use serde_json::{json, to_value, Value}; use serde_json::{json, to_value, Value};
use tokio::sync::mpsc; use tokio::sync::mpsc;
@@ -35,10 +33,11 @@ use super::{
}, },
responses::{mask_tools_as_mcp, patch_streaming_response_json}, responses::{mask_tools_as_mcp, patch_streaming_response_json},
streaming::handle_streaming_response, streaming::handle_streaming_response,
utils::{apply_provider_headers, extract_auth_header, probe_endpoint_for_model}, utils::{apply_provider_headers, extract_auth_header},
}; };
use crate::{ use crate::{
core::{CircuitBreaker, CircuitBreakerConfig as CoreCircuitBreakerConfig}, app_context::AppContext,
core::{ModelCard, RuntimeType, Worker, WorkerRegistry},
data_connector::{ data_connector::{
ConversationId, ConversationItemStorage, ConversationStorage, ListParams, ResponseId, ConversationId, ConversationItemStorage, ConversationStorage, ListParams, ResponseId,
ResponseStorage, SortOrder, ResponseStorage, SortOrder,
@@ -56,7 +55,6 @@ use crate::{
ResponsesGetParams, ResponsesRequest, ResponsesGetParams, ResponsesRequest,
}, },
}, },
routers::header_utils::apply_request_headers,
}; };
// ============================================================================ // ============================================================================
@@ -89,23 +87,16 @@ static SGLANG_FIELDS: Lazy<HashSet<&'static str>> = Lazy::new(|| {
]) ])
}); });
/// Cached endpoint information
#[derive(Clone, Debug)]
struct CachedEndpoint {
url: String,
cached_at: Instant,
}
/// Router for OpenAI backend /// 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 /// HTTP client for upstream OpenAI-compatible API
client: reqwest::Client, client: reqwest::Client,
/// Multiple OpenAI-compatible API endpoints (OpenAI, xAI, etc.) /// Worker registry for model-based worker lookup
worker_urls: Vec<String>, worker_registry: Arc<WorkerRegistry>,
/// Model cache: model_id -> endpoint URL
model_cache: Arc<DashMap<String, CachedEndpoint>>,
/// Circuit breaker
circuit_breaker: CircuitBreaker,
/// Health status /// Health status
healthy: AtomicBool, healthy: AtomicBool,
/// Response storage for managing conversation history /// Response storage for managing conversation history
@@ -120,8 +111,11 @@ pub struct OpenAIRouter {
impl std::fmt::Debug for OpenAIRouter { impl std::fmt::Debug for OpenAIRouter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let registry_stats = self.worker_registry.stats();
f.debug_struct("OpenAIRouter") f.debug_struct("OpenAIRouter")
.field("worker_urls", &self.worker_urls) .field("registered_workers", &registry_stats.total_workers)
.field("registered_models", &registry_stats.total_models)
.field("healthy_workers", &registry_stats.healthy_workers)
.field("healthy", &self.healthy) .field("healthy", &self.healthy)
.finish() .finish()
} }
@@ -131,33 +125,16 @@ impl OpenAIRouter {
/// Maximum number of conversation items to attach as input when a conversation is provided /// 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;
/// Model discovery cache TTL (1 hour)
const MODEL_CACHE_TTL_SECS: u64 = 3600;
/// Create a new OpenAI router /// Create a new OpenAI router
pub async fn new( ///
worker_urls: Vec<String>, /// Workers are registered separately via the external worker registration workflow.
ctx: &Arc<crate::app_context::AppContext>, /// This router queries the WorkerRegistry to find workers that support requested models.
) -> Result<Self, String> { pub async fn new(ctx: &Arc<AppContext>) -> Result<Self, String> {
// Use HTTP client from AppContext // Use HTTP client from AppContext
let client = ctx.client.clone(); let client = ctx.client.clone();
// Normalize URLs (remove trailing slashes) // Get worker registry from AppContext
let worker_urls: Vec<String> = worker_urls let worker_registry = ctx.worker_registry.clone();
.into_iter()
.map(|url| url.trim_end_matches('/').to_string())
.collect();
// Convert circuit breaker config from AppContext
let cb = &ctx.router_config.circuit_breaker;
let core_cb_config = CoreCircuitBreakerConfig {
failure_threshold: cb.failure_threshold,
success_threshold: cb.success_threshold,
timeout_duration: Duration::from_secs(cb.timeout_duration_secs),
window_duration: Duration::from_secs(cb.window_duration_secs),
};
let circuit_breaker = CircuitBreaker::with_config(core_cb_config);
// Get MCP manager from AppContext (must be initialized) // Get MCP manager from AppContext (must be initialized)
let mcp_manager = ctx let mcp_manager = ctx
@@ -168,9 +145,7 @@ impl OpenAIRouter {
Ok(Self { Ok(Self {
client, client,
worker_urls, worker_registry,
model_cache: Arc::new(DashMap::new()),
circuit_breaker,
healthy: AtomicBool::new(true), healthy: AtomicBool::new(true),
response_storage: ctx.response_storage.clone(), response_storage: ctx.response_storage.clone(),
conversation_storage: ctx.conversation_storage.clone(), conversation_storage: ctx.conversation_storage.clone(),
@@ -179,76 +154,178 @@ impl OpenAIRouter {
}) })
} }
/// Discover which endpoint has the model /// Refresh models for a single external worker by querying its /v1/models endpoint.
async fn find_endpoint_for_model( ///
/// Returns true if refresh succeeded and models were cached on the worker.
async fn refresh_worker_models(
&self,
worker: &Arc<dyn Worker>,
auth_header: Option<&HeaderValue>,
) -> bool {
let url = format!("{}/v1/models", worker.url());
// Build request to backend
let mut backend_req = self.client.get(&url);
if let Some(auth) = auth_header {
backend_req = apply_provider_headers(backend_req, &url, Some(auth));
}
match backend_req.send().await {
Ok(response) if response.status().is_success() => {
match response.json::<Value>().await {
Ok(json_response) => {
if let Some(data) = json_response.get("data").and_then(|d| d.as_array()) {
let model_cards: Vec<ModelCard> = data
.iter()
.filter_map(|m| m.get("id").and_then(|id| id.as_str()))
.map(ModelCard::new)
.collect();
if !model_cards.is_empty() {
tracing::info!(
"Model refresh: found {} models from {}",
model_cards.len(),
url
);
worker.set_models(model_cards);
return true;
}
}
false
}
Err(e) => {
tracing::warn!("Failed to parse models response: {}", e);
false
}
}
}
Ok(response) => {
tracing::debug!(
"Model refresh returned non-success status {} from {}",
response.status(),
url
);
false
}
Err(e) => {
tracing::warn!("Failed to fetch models from backend: {}", e);
false
}
}
}
/// Refresh models for ALL external workers in parallel.
async fn refresh_external_models(&self, auth_header: Option<&HeaderValue>) {
let external_workers = self.worker_registry.get_workers_filtered(
None,
None,
None,
Some(RuntimeType::External),
true, // healthy_only
);
if external_workers.is_empty() {
return;
}
tracing::debug!(
"Refreshing models for {} external workers",
external_workers.len()
);
// Refresh all workers in parallel
let futures: Vec<_> = external_workers
.iter()
.map(|w| self.refresh_worker_models(w, auth_header))
.collect();
join_all(futures).await;
}
/// Select a worker for the given model using the WorkerRegistry.
///
/// This method queries the registry for external workers (RuntimeType::External)
/// that support the requested model. It checks:
/// 1. Workers registered with matching model ID (including aliases via ModelCard)
/// 2. Worker health status
/// 3. Circuit breaker state
///
/// If no worker is found with explicit model support, it will refresh models
/// on all external workers in parallel, then retry the search.
///
/// Returns an error response if no suitable worker is found.
async fn select_worker_for_model(
&self, &self,
model_id: &str, model_id: &str,
auth_header: Option<&str>, auth_header: Option<&HeaderValue>,
) -> Result<String, Response> { ) -> Result<Arc<dyn Worker>, Box<Response>> {
// Single endpoint - fast path // Helper to find candidates for a model
if self.worker_urls.len() == 1 { // Note: We get ALL external workers and filter by supports_model() because
return Ok(self.worker_urls[0].clone()); // wildcard workers (empty models) aren't in the model index but support any model
let find_candidates = || {
self.worker_registry
.get_workers_filtered(
None, // Get all external workers, not just those in model index
None,
None,
Some(RuntimeType::External),
true, // healthy_only
)
.into_iter()
.filter(|w| w.supports_model(model_id) && w.circuit_breaker().can_execute())
.collect::<Vec<_>>()
};
// First try: find workers that already support this model
let candidates = find_candidates();
if !candidates.is_empty() {
return Ok(candidates
.into_iter()
.min_by_key(|w| w.load())
.expect("candidates is not empty"));
} }
// Check cache // No match found - refresh models on all external workers
if let Some(entry) = self.model_cache.get(model_id) { tracing::debug!(
if entry.cached_at.elapsed() < Duration::from_secs(Self::MODEL_CACHE_TTL_SECS) { "No worker found for model '{}', refreshing external worker models",
return Ok(entry.url.clone()); model_id
}
}
// Probe all endpoints in parallel
let mut handles = vec![];
let model = model_id.to_string();
let auth = auth_header.map(|s| s.to_string());
for url in &self.worker_urls {
let handle = tokio::spawn(probe_endpoint_for_model(
self.client.clone(),
url.clone(),
model.clone(),
auth.clone(),
));
handles.push(handle);
}
// Return first successful endpoint
for handle in handles {
if let Ok(Ok(url)) = handle.await {
// Cache it
self.model_cache.insert(
model_id.to_string(),
CachedEndpoint {
url: url.clone(),
cached_at: Instant::now(),
},
); );
return Ok(url); self.refresh_external_models(auth_header).await;
}
// Second try: check if any worker now supports the model after refresh
let candidates = find_candidates();
if !candidates.is_empty() {
return Ok(candidates
.into_iter()
.min_by_key(|w| w.load())
.expect("candidates is not empty"));
} }
// Model not found on any endpoint Err(Box::new(
Err(( (
StatusCode::NOT_FOUND, StatusCode::NOT_FOUND,
Json(json!({ Json(json!({
"error": { "error": {
"message": format!("Model '{}' not found on any endpoint", model_id), "message": format!("No worker available for model '{}'", model_id),
"type": "model_not_found", "type": "model_not_found",
} }
})), })),
) )
.into_response()) .into_response(),
))
} }
/// Handle non-streaming response with optional MCP tool loop /// Handle non-streaming response with optional MCP tool loop
async fn handle_non_streaming_response( async fn handle_non_streaming_response(
&self, &self,
url: String, worker: &Arc<dyn Worker>,
headers: Option<&HeaderMap>, headers: Option<&HeaderMap>,
mut payload: Value, mut payload: Value,
original_body: &ResponsesRequest, original_body: &ResponsesRequest,
original_previous_response_id: Option<String>, original_previous_response_id: Option<String>,
) -> Response { ) -> Response {
let url = format!("{}/v1/responses", worker.url());
// Check if MCP is active for this request // Check if MCP is active for this request
// Ensure dynamic client is created if needed // Ensure dynamic client is created if needed
if let Some(ref tools) = original_body.tools { if let Some(ref tools) = original_body.tools {
@@ -284,7 +361,7 @@ impl OpenAIRouter {
{ {
Ok(resp) => response_json = resp, Ok(resp) => response_json = resp,
Err(err) => { Err(err) => {
self.circuit_breaker.record_failure(); worker.circuit_breaker().record_failure();
return ( return (
StatusCode::INTERNAL_SERVER_ERROR, StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({"error": {"message": err}})), Json(json!({"error": {"message": err}})),
@@ -296,14 +373,16 @@ impl OpenAIRouter {
// No MCP - simple request // No MCP - simple request
let mut request_builder = self.client.post(&url).json(&payload); let mut request_builder = self.client.post(&url).json(&payload);
if let Some(h) = headers {
request_builder = apply_request_headers(h, request_builder, true); // 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());
let response = match request_builder.send().await { let response = match request_builder.send().await {
Ok(r) => r, Ok(r) => r,
Err(e) => { Err(e) => {
self.circuit_breaker.record_failure(); worker.circuit_breaker().record_failure();
tracing::error!( tracing::error!(
url = %url, url = %url,
error = %e, error = %e,
@@ -318,7 +397,7 @@ impl OpenAIRouter {
}; };
if !response.status().is_success() { if !response.status().is_success() {
self.circuit_breaker.record_failure(); worker.circuit_breaker().record_failure();
let status = StatusCode::from_u16(response.status().as_u16()) let status = StatusCode::from_u16(response.status().as_u16())
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
let body = response.text().await.unwrap_or_default(); let body = response.text().await.unwrap_or_default();
@@ -328,7 +407,7 @@ impl OpenAIRouter {
response_json = match response.json::<Value>().await { response_json = match response.json::<Value>().await {
Ok(r) => r, Ok(r) => r,
Err(e) => { Err(e) => {
self.circuit_breaker.record_failure(); worker.circuit_breaker().record_failure();
return ( return (
StatusCode::INTERNAL_SERVER_ERROR, StatusCode::INTERNAL_SERVER_ERROR,
format!("Failed to parse upstream response: {}", e), format!("Failed to parse upstream response: {}", e),
@@ -337,7 +416,7 @@ impl OpenAIRouter {
} }
}; };
self.circuit_breaker.record_success(); worker.circuit_breaker().record_success();
} }
// Patch response with metadata // Patch response with metadata
@@ -376,134 +455,134 @@ 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 all endpoints in parallel - only healthy if ALL are healthy // Check health of all external workers
if self.worker_urls.is_empty() { let external_workers: Vec<_> = self
return (StatusCode::SERVICE_UNAVAILABLE, "No endpoints configured").into_response(); .worker_registry
.get_all()
.into_iter()
.filter(|w| w.metadata().runtime_type == RuntimeType::External)
.collect();
if external_workers.is_empty() {
return (
StatusCode::SERVICE_UNAVAILABLE,
"No external workers registered",
)
.into_response();
} }
let mut handles = vec![]; let mut healthy_count = 0;
for url in &self.worker_urls { let mut unhealthy_workers = Vec::new();
let url = url.clone();
let client = self.client.clone();
let handle = tokio::spawn(async move { for worker in &external_workers {
let probe_url = format!("{}/v1/models", url); if worker.is_healthy() {
match client healthy_count += 1;
.get(&probe_url)
.timeout(Duration::from_secs(2))
.send()
.await
{
Ok(resp) => {
let code = resp.status();
// Treat success and auth-required as healthy (endpoint reachable)
if code.is_success() || code.as_u16() == 401 || code.as_u16() == 403 {
Ok(())
} else { } else {
Err(format!("Endpoint {} returned status {}", url, code)) unhealthy_workers.push(format!("{} ({})", worker.model_id(), worker.url()));
}
}
Err(e) => Err(format!("Endpoint {} error: {}", url, e)),
}
});
handles.push(handle);
}
// Collect all results
let mut errors = Vec::new();
for handle in handles {
match handle.await {
Ok(Ok(())) => (),
Ok(Err(e)) => errors.push(e),
Err(e) => errors.push(format!("Task join error: {}", e)),
} }
} }
if errors.is_empty() { if unhealthy_workers.is_empty() {
(StatusCode::OK, "OK").into_response() (
StatusCode::OK,
format!("OK - {} workers healthy", healthy_count),
)
.into_response()
} else { } else {
( (
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
format!("Some endpoints unhealthy: {}", errors.join(", ")), format!(
"{}/{} workers unhealthy: {}",
unhealthy_workers.len(),
external_workers.len(),
unhealthy_workers.join(", ")
),
) )
.into_response() .into_response()
} }
} }
async fn get_server_info(&self, _req: Request<Body>) -> Response { async fn get_server_info(&self, _req: Request<Body>) -> Response {
let stats = self.worker_registry.stats();
let external_workers: Vec<_> = self
.worker_registry
.get_all()
.into_iter()
.filter(|w| w.metadata().runtime_type == RuntimeType::External)
.collect();
let worker_urls: Vec<String> = external_workers
.iter()
.map(|w| w.url().to_string())
.collect();
let info = json!({ let info = json!({
"router_type": "openai", "router_type": "openai",
"workers": self.worker_urls.len(), "total_workers": stats.total_workers,
"worker_urls": &self.worker_urls "external_workers": external_workers.len(),
"healthy_workers": stats.healthy_workers,
"total_models": stats.total_models,
"worker_urls": worker_urls
}); });
(StatusCode::OK, info.to_string()).into_response() (StatusCode::OK, info.to_string()).into_response()
} }
async fn get_models(&self, req: Request<Body>) -> Response { async fn get_models(&self, req: Request<Body>) -> Response {
// Aggregate models from all endpoints // Return models from all registered external workers
if self.worker_urls.is_empty() { let external_workers: Vec<_> = self
return (StatusCode::SERVICE_UNAVAILABLE, "No endpoints configured").into_response(); .worker_registry
.get_all()
.into_iter()
.filter(|w| w.metadata().runtime_type == RuntimeType::External)
.collect();
if external_workers.is_empty() {
return (
StatusCode::SERVICE_UNAVAILABLE,
"No external workers registered",
)
.into_response();
} }
let headers = req.headers(); // Refresh models for all external workers using user's auth header
let auth = headers let auth_header = extract_auth_header(Some(req.headers()), &None);
.get("authorization") self.refresh_external_models(auth_header.as_ref()).await;
.or_else(|| headers.get("Authorization"));
// Query all endpoints in parallel // Collect models from all workers
let mut handles = vec![];
for url in &self.worker_urls {
let url = url.clone();
let client = self.client.clone();
let auth = auth.cloned();
let handle = tokio::spawn(async move {
let models_url = format!("{}/v1/models", url);
let req = client.get(&models_url);
// Apply provider-specific headers (handles Anthropic, xAI, OpenAI, etc.)
let req = apply_provider_headers(req, &url, auth.as_ref());
match req.send().await {
Ok(res) => {
if res.status().is_success() {
match res.json::<Value>().await {
Ok(json) => Ok(json),
Err(e) => {
tracing::warn!(
"Failed to parse models response from '{}': {}",
url,
e
);
Err(())
}
}
} else {
tracing::warn!(
"Getting models from '{}' failed with status: {}",
url,
res.status()
);
Err(())
}
}
Err(e) => {
tracing::warn!("Request to get models from '{}' failed: {}", url, e);
Err(())
}
}
});
handles.push(handle);
}
// Collect all model lists
let mut all_models = Vec::new(); let mut all_models = Vec::new();
for handle in handles { let mut seen_models = HashSet::new();
if let Ok(Ok(json)) = handle.await {
if let Some(data) = json.get("data").and_then(|v| v.as_array()) { for worker in &external_workers {
all_models.extend_from_slice(data); for model_card in worker.models() {
let owned_by = model_card
.provider
.as_ref()
.map(|p| format!("{:?}", p).to_lowercase())
.unwrap_or_else(|| "unknown".to_string());
// Add primary model ID
if seen_models.insert(model_card.id.clone()) {
all_models.push(json!({
"id": &model_card.id,
"object": "model",
"created": 0,
"owned_by": &owned_by,
"aliases": model_card.aliases,
"model_type": format!("{:?}", model_card.model_type),
}));
}
// Add aliases as separate entries for compatibility
for alias in &model_card.aliases {
if seen_models.insert(alias.clone()) {
all_models.push(json!({
"id": alias,
"object": "model",
"created": 0,
"owned_by": &owned_by,
"primary_model": &model_card.id,
}));
}
} }
} }
} }
@@ -546,20 +625,16 @@ impl crate::routers::RouterTrait for OpenAIRouter {
body: &ChatCompletionRequest, body: &ChatCompletionRequest,
_model_id: Option<&str>, _model_id: Option<&str>,
) -> Response { ) -> Response {
if !self.circuit_breaker.can_execute() { // Extract auth header for passthrough mode
return (StatusCode::SERVICE_UNAVAILABLE, "Circuit breaker open").into_response(); let auth_header = extract_auth_header(headers, &None);
}
// Extract auth header // Select worker for model (discovery happens inside if needed)
let auth = extract_auth_header(headers); let worker = match self
.select_worker_for_model(body.model.as_str(), auth_header.as_ref())
// Find endpoint for model
let base_url = match self
.find_endpoint_for_model(body.model.as_str(), auth)
.await .await
{ {
Ok(url) => url, Ok(w) => w,
Err(response) => return response, Err(response) => return *response,
}; };
// Serialize request body, removing SGLang-only fields // Serialize request body, removing SGLang-only fields
@@ -582,15 +657,13 @@ impl crate::routers::RouterTrait for OpenAIRouter {
} }
} }
let url = format!("{}/v1/chat/completions", base_url); let url = format!("{}/v1/chat/completions", worker.url());
let mut req = self.client.post(&url).json(&payload); let mut req = self.client.post(&url).json(&payload);
// Forward Authorization header if provided // Apply provider-specific headers (handles Anthropic x-api-key, etc.)
if let Some(h) = headers { // Passthrough mode: user's auth header takes priority, worker's key is fallback
if let Some(auth) = h.get("authorization").or_else(|| h.get("Authorization")) { let auth_header = extract_auth_header(headers, worker.api_key());
req = req.header("Authorization", auth); req = apply_provider_headers(req, &url, auth_header.as_ref());
}
}
// Accept SSE when stream=true // Accept SSE when stream=true
if body.stream { if body.stream {
@@ -600,7 +673,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
let resp = match req.send().await { let resp = match req.send().await {
Ok(r) => r, Ok(r) => r,
Err(e) => { Err(e) => {
self.circuit_breaker.record_failure(); worker.circuit_breaker().record_failure();
return ( return (
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
format!("Failed to contact upstream: {}", e), format!("Failed to contact upstream: {}", e),
@@ -617,7 +690,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
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) => {
self.circuit_breaker.record_success(); worker.circuit_breaker().record_success();
let mut response = Response::new(Body::from(body)); let mut response = Response::new(Body::from(body));
*response.status_mut() = status; *response.status_mut() = status;
if let Some(ct) = content_type { if let Some(ct) = content_type {
@@ -626,7 +699,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
response response
} }
Err(e) => { Err(e) => {
self.circuit_breaker.record_failure(); worker.circuit_breaker().record_failure();
( (
StatusCode::INTERNAL_SERVER_ERROR, StatusCode::INTERNAL_SERVER_ERROR,
format!("Failed to read response: {}", e), format!("Failed to read response: {}", e),
@@ -683,18 +756,19 @@ impl crate::routers::RouterTrait for OpenAIRouter {
body: &ResponsesRequest, body: &ResponsesRequest,
model_id: Option<&str>, model_id: Option<&str>,
) -> Response { ) -> Response {
// Extract auth header // Extract auth header for passthrough mode
let auth = extract_auth_header(headers); let auth_header = extract_auth_header(headers, &None);
// Find endpoint for model (use model_id if provided, otherwise use body.model) // 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 base_url = match self.find_endpoint_for_model(model, auth).await { let worker = match self
Ok(url) => url, .select_worker_for_model(model, auth_header.as_ref())
Err(response) => return response, .await
{
Ok(w) => w,
Err(response) => return *response,
}; };
let url = format!("{}/v1/responses", base_url);
// Clone the body for validation and logic, but we'll build payload differently // 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 {
@@ -992,10 +1066,11 @@ impl crate::routers::RouterTrait for OpenAIRouter {
} }
// Delegate to streaming or non-streaming handler // Delegate to streaming or non-streaming handler
let url = format!("{}/v1/responses", worker.url());
if body.stream.unwrap_or(false) { if body.stream.unwrap_or(false) {
handle_streaming_response( handle_streaming_response(
&self.client, &self.client,
&self.circuit_breaker, worker.circuit_breaker(),
Some(&self.mcp_manager), Some(&self.mcp_manager),
self.response_storage.clone(), self.response_storage.clone(),
self.conversation_storage.clone(), self.conversation_storage.clone(),
@@ -1009,7 +1084,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.await .await
} else { } else {
self.handle_non_streaming_response( self.handle_non_streaming_response(
url, &worker,
headers, headers,
payload, payload,
body, body,
+28 -59
View File
@@ -2,7 +2,7 @@
use std::collections::HashMap; use std::collections::HashMap;
use axum::http::{HeaderMap, HeaderValue}; use axum::http::HeaderValue;
// ============================================================================ // ============================================================================
// SSE Event Type Constants // SSE Event Type Constants
@@ -99,16 +99,6 @@ impl OutputIndexMapper {
// Provider Detection and Header Handling // Provider Detection and Header Handling
// ============================================================================ // ============================================================================
/// Extract authorization header from request headers
/// Checks both "authorization" and "Authorization" (case variations)
pub fn extract_auth_header(headers: Option<&HeaderMap>) -> Option<&str> {
headers.and_then(|h| {
h.get("authorization")
.or_else(|| h.get("Authorization"))
.and_then(|v| v.to_str().ok())
})
}
/// API provider types /// API provider types
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ApiProvider { pub enum ApiProvider {
@@ -168,56 +158,35 @@ pub fn apply_provider_headers(
req req
} }
/// Probe a single endpoint to check if it has the model // ============================================================================
/// Returns Ok(url) if model found, Err(()) otherwise // Auth Header Resolution
pub async fn probe_endpoint_for_model( // ============================================================================
client: reqwest::Client,
url: String,
model: String,
auth: Option<String>,
) -> Result<String, ()> {
use tracing::debug;
let probe_url = format!("{}/v1/models/{}", url, model); /// Extract auth header with passthrough semantics.
let req = client ///
.get(&probe_url) /// Passthrough mode: User's Authorization header takes priority.
.timeout(std::time::Duration::from_secs(5)); /// Fallback: Worker's API key is used only if user didn't provide auth.
///
/// This enables use cases where:
/// 1. Users send their own API keys (multi-tenant, BYOK)
/// 2. Router has a default key for users who don't provide one
pub fn extract_auth_header(
headers: Option<&http::HeaderMap>,
worker_api_key: &Option<String>,
) -> Option<HeaderValue> {
// Passthrough: Try user's auth header first
let user_auth = headers.and_then(|h| {
h.get("authorization")
.or_else(|| h.get("Authorization"))
.cloned()
});
// Apply provider-specific headers (handles Anthropic, xAI, OpenAI, etc.) // Return user's auth if provided, otherwise use worker's API key
let auth_header_value = auth.as_ref().and_then(|a| HeaderValue::from_str(a).ok()); user_auth.or_else(|| {
let req = apply_provider_headers(req, &url, auth_header_value.as_ref()); worker_api_key
.as_ref()
match req.send().await { .and_then(|k| HeaderValue::from_str(&format!("Bearer {}", k)).ok())
Ok(resp) => { })
let status = resp.status();
if status.is_success() {
debug!(
url = %url,
model = %model,
status = %status,
"Model found on endpoint"
);
Ok(url)
} else {
debug!(
url = %url,
model = %model,
status = %status,
"Model not found on endpoint (unsuccessful status)"
);
Err(())
}
}
Err(e) => {
debug!(
url = %url,
model = %model,
error = %e,
"Probe request to endpoint failed"
);
Err(())
}
}
} }
// ============================================================================ // ============================================================================
+44 -2
View File
@@ -16,8 +16,10 @@ use std::{
use serde_json::json; use serde_json::json;
use sgl_model_gateway::{ use sgl_model_gateway::{
app_context::AppContext, app_context::AppContext,
config::RouterConfig, config::{RouterConfig, RoutingMode},
core::{LoadMonitor, WorkerRegistry}, core::{
BasicWorkerBuilder, LoadMonitor, ModelCard, RuntimeType, Worker, WorkerRegistry, WorkerType,
},
data_connector::{ data_connector::{
MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage, MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage,
}, },
@@ -111,6 +113,26 @@ pub async fn create_test_context(config: RouterConfig) -> Arc<AppContext> {
.set(engine) .set(engine)
.expect("WorkflowEngine should only be initialized once"); .expect("WorkflowEngine should only be initialized once");
// Register external workers for OpenAI mode
if let RoutingMode::OpenAI { worker_urls, .. } = &config.mode {
for url in worker_urls {
// Create a worker that supports common test models
let models = vec![
ModelCard::new("mock-model"),
ModelCard::new("gpt-4"),
ModelCard::new("gpt-3.5-turbo"),
];
let worker: Arc<dyn Worker> = Arc::new(
BasicWorkerBuilder::new(url)
.worker_type(WorkerType::Regular)
.runtime_type(RuntimeType::External)
.models(models)
.build(),
);
app_context.worker_registry.register(worker);
}
}
// Initialize MCP manager with empty config // Initialize MCP manager with empty config
use sgl_model_gateway::mcp::{McpConfig, McpManager}; use sgl_model_gateway::mcp::{McpConfig, McpManager};
let empty_config = McpConfig { let empty_config = McpConfig {
@@ -222,6 +244,26 @@ pub async fn create_test_context_with_mcp_config(
.set(engine) .set(engine)
.expect("WorkflowEngine should only be initialized once"); .expect("WorkflowEngine should only be initialized once");
// Register external workers for OpenAI mode
if let RoutingMode::OpenAI { worker_urls, .. } = &config.mode {
for url in worker_urls {
// Create a worker that supports common test models
let models = vec![
ModelCard::new("mock-model"),
ModelCard::new("gpt-4"),
ModelCard::new("gpt-3.5-turbo"),
];
let worker: Arc<dyn Worker> = Arc::new(
BasicWorkerBuilder::new(url)
.worker_type(WorkerType::Regular)
.runtime_type(RuntimeType::External)
.models(models)
.build(),
);
app_context.worker_registry.register(worker);
}
}
// Initialize MCP manager from config file // Initialize MCP manager from config file
let mcp_config = McpConfig::from_file(mcp_config_path) let mcp_config = McpConfig::from_file(mcp_config_path)
.await .await
+50 -1
View File
@@ -5,7 +5,9 @@ use reqwest::Client;
use sgl_model_gateway::{ use sgl_model_gateway::{
app_context::AppContext, app_context::AppContext,
config::RouterConfig, config::RouterConfig,
core::{LoadMonitor, WorkerRegistry}, core::{
BasicWorkerBuilder, LoadMonitor, ModelCard, RuntimeType, Worker, WorkerRegistry, WorkerType,
},
data_connector::{ data_connector::{
MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage, MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage,
}, },
@@ -209,3 +211,50 @@ pub async fn create_test_app_context() -> Arc<AppContext> {
.unwrap(), .unwrap(),
) )
} }
/// Register an external worker (OpenAI-compatible API endpoint) in the test AppContext.
///
/// This is used by tests that need to test the OpenAI router, which expects
/// workers to be registered in the WorkerRegistry before routing requests.
///
/// # Arguments
/// * `ctx` - The AppContext to register the worker in
/// * `url` - The base URL of the external API endpoint
/// * `models` - Optional list of model IDs this worker supports. If empty, uses "gpt-3.5-turbo" as default.
#[allow(dead_code)]
pub fn register_external_worker(ctx: &Arc<AppContext>, url: &str, models: Option<Vec<&str>>) {
let model_list: Vec<ModelCard> = models
.unwrap_or_else(|| vec!["gpt-3.5-turbo"])
.into_iter()
.map(ModelCard::new)
.collect();
let worker: Arc<dyn Worker> = Arc::new(
BasicWorkerBuilder::new(url)
.worker_type(WorkerType::Regular)
.runtime_type(RuntimeType::External)
.models(model_list)
.build(),
);
ctx.worker_registry.register(worker);
}
/// Register an external worker with a custom model card that has aliases.
///
/// # Arguments
/// * `ctx` - The AppContext to register the worker in
/// * `url` - The base URL of the external API endpoint
/// * `model_card` - A fully configured ModelCard with aliases, provider, etc.
#[allow(dead_code)]
pub fn register_external_worker_with_card(ctx: &Arc<AppContext>, url: &str, model_card: ModelCard) {
let worker: Arc<dyn Worker> = Arc::new(
BasicWorkerBuilder::new(url)
.worker_type(WorkerType::Regular)
.runtime_type(RuntimeType::External)
.model(model_card)
.build(),
);
ctx.worker_registry.register(worker);
}
+37 -47
View File
@@ -96,7 +96,9 @@ fn create_minimal_completion_request() -> CompletionRequest {
#[tokio::test] #[tokio::test]
async fn test_openai_router_creation() { async fn test_openai_router_creation() {
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec!["https://api.openai.com".to_string()], &ctx).await; // Register an external worker before creating the router
common::test_app::register_external_worker(&ctx, "https://api.openai.com", None);
let router = OpenAIRouter::new(&ctx).await;
assert!(router.is_ok(), "Router creation should succeed"); assert!(router.is_ok(), "Router creation should succeed");
@@ -109,9 +111,8 @@ async fn test_openai_router_creation() {
#[tokio::test] #[tokio::test]
async fn test_openai_router_server_info() { async fn test_openai_router_server_info() {
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec!["https://api.openai.com".to_string()], &ctx) common::test_app::register_external_worker(&ctx, "https://api.openai.com", None);
.await let router = OpenAIRouter::new(&ctx).await.unwrap();
.unwrap();
let req = Request::builder() let req = Request::builder()
.method(Method::GET) .method(Method::GET)
@@ -135,9 +136,8 @@ async fn test_openai_router_models() {
// Use mock server for deterministic models response // Use mock server for deterministic models response
let mock_server = MockOpenAIServer::new().await; let mock_server = MockOpenAIServer::new().await;
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec![mock_server.base_url()], &ctx) common::test_app::register_external_worker(&ctx, &mock_server.base_url(), None);
.await let router = OpenAIRouter::new(&ctx).await.unwrap();
.unwrap();
let req = Request::builder() let req = Request::builder()
.method(Method::GET) .method(Method::GET)
@@ -209,7 +209,8 @@ async fn test_openai_router_responses_with_mock() {
let base_url = format!("http://{}", addr); let base_url = format!("http://{}", addr);
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); common::test_app::register_external_worker(&ctx, &base_url, Some(vec!["gpt-4o-mini"]));
let router = OpenAIRouter::new(&ctx).await.unwrap();
// Get storage from context (router uses this, not a separate storage) // Get storage from context (router uses this, not a separate storage)
let storage = ctx.response_storage.clone(); let storage = ctx.response_storage.clone();
@@ -473,7 +474,8 @@ async fn test_openai_router_responses_streaming_with_mock() {
let base_url = format!("http://{}", addr); let base_url = format!("http://{}", addr);
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); common::test_app::register_external_worker(&ctx, &base_url, Some(vec!["gpt-5-nano"]));
let router = OpenAIRouter::new(&ctx).await.unwrap();
// Get storage from context and seed a previous response // Get storage from context and seed a previous response
let storage = ctx.response_storage.clone(); let storage = ctx.response_storage.clone();
@@ -598,9 +600,8 @@ async fn test_router_factory_openai_mode() {
#[tokio::test] #[tokio::test]
async fn test_unsupported_endpoints() { async fn test_unsupported_endpoints() {
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec!["https://api.openai.com".to_string()], &ctx) common::test_app::register_external_worker(&ctx, "https://api.openai.com", None);
.await let router = OpenAIRouter::new(&ctx).await.unwrap();
.unwrap();
let generate_request = GenerateRequest { let generate_request = GenerateRequest {
text: Some("Hello world".to_string()), text: Some("Hello world".to_string()),
@@ -658,8 +659,9 @@ async fn test_openai_router_chat_completion_with_mock() {
let base_url = mock_server.base_url(); let base_url = mock_server.base_url();
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
// Create router pointing to mock server // Register the mock server worker and create router
let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); common::test_app::register_external_worker(&ctx, &base_url, None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
// Create a minimal chat completion request // Create a minimal chat completion request
let mut chat_request = create_minimal_chat_request(); let mut chat_request = create_minimal_chat_request();
@@ -693,8 +695,9 @@ async fn test_openai_e2e_with_server() {
let base_url = mock_server.base_url(); let base_url = mock_server.base_url();
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
// Create router // Register the mock server worker and create router
let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); common::test_app::register_external_worker(&ctx, &base_url, None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
// Create Axum app with chat completions endpoint // Create Axum app with chat completions endpoint
let app = Router::new().route( let app = Router::new().route(
@@ -758,7 +761,8 @@ async fn test_openai_router_chat_streaming_with_mock() {
let mock_server = MockOpenAIServer::new().await; let mock_server = MockOpenAIServer::new().await;
let base_url = mock_server.base_url(); let base_url = mock_server.base_url();
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); common::test_app::register_external_worker(&ctx, &base_url, None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
// Build a streaming chat request // Build a streaming chat request
let val = json!({ let val = json!({
@@ -797,9 +801,8 @@ async fn test_openai_router_chat_streaming_with_mock() {
#[tokio::test] #[tokio::test]
async fn test_openai_router_circuit_breaker() { async fn test_openai_router_circuit_breaker() {
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec!["http://invalid-url-that-will-fail".to_string()], &ctx) common::test_app::register_external_worker(&ctx, "http://invalid-url-that-will-fail", None);
.await let router = OpenAIRouter::new(&ctx).await.unwrap();
.unwrap();
let chat_request = create_minimal_chat_request(); let chat_request = create_minimal_chat_request();
@@ -814,19 +817,19 @@ async fn test_openai_router_circuit_breaker() {
} }
} }
/// Test that Authorization header is forwarded in /v1/models /// Test that /v1/models returns models from registered workers' ModelCards
///
/// With the new worker-based design, models are returned from the WorkerRegistry
/// and don't require calling external APIs. Auth headers are used for routing
/// requests to workers, not for the models endpoint.
#[tokio::test] #[tokio::test]
async fn test_openai_router_models_auth_forwarding() { async fn test_openai_router_models_from_registry() {
// Start a mock server that requires Authorization
let expected_auth = "Bearer test-token".to_string();
let mock_server = MockOpenAIServer::new_with_auth(Some(expected_auth.clone())).await;
let ctx = common::test_app::create_test_app_context().await; let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec![mock_server.base_url()], &ctx) // Register a worker with the default model
.await common::test_app::register_external_worker(&ctx, "https://api.example.com", None);
.unwrap(); let router = OpenAIRouter::new(&ctx).await.unwrap();
// 1) Without auth header -> expect 200 with empty model list // Get models - should return the registered model
// (multi-endpoint aggregation silently skips failed endpoints)
let req = Request::builder() let req = Request::builder()
.method(Method::GET) .method(Method::GET)
.uri("/models") .uri("/models")
@@ -840,24 +843,11 @@ async fn test_openai_router_models_auth_forwarding() {
let body_str = String::from_utf8(body_bytes.to_vec()).unwrap(); let body_str = String::from_utf8(body_bytes.to_vec()).unwrap();
let models: serde_json::Value = serde_json::from_str(&body_str).unwrap(); let models: serde_json::Value = serde_json::from_str(&body_str).unwrap();
assert_eq!(models["object"], "list"); assert_eq!(models["object"], "list");
assert_eq!(models["data"].as_array().unwrap().len(), 0); // Empty when auth fails
// 2) With auth header -> expect 200 // Should have the default model (gpt-3.5-turbo)
let req = Request::builder() let data = models["data"].as_array().unwrap();
.method(Method::GET) assert_eq!(data.len(), 1);
.uri("/models") assert_eq!(data[0]["id"], "gpt-3.5-turbo");
.header("Authorization", expected_auth)
.body(Body::empty())
.unwrap();
let response = router.get_models(req).await;
assert_eq!(response.status(), StatusCode::OK);
let (_, body) = response.into_parts();
let body_bytes = axum::body::to_bytes(body, usize::MAX).await.unwrap();
let body_str = String::from_utf8(body_bytes.to_vec()).unwrap();
let models: serde_json::Value = serde_json::from_str(&body_str).unwrap();
assert_eq!(models["object"], "list");
} }
#[test] #[test]