[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,
sync::{
atomic::{AtomicBool, AtomicUsize, Ordering},
Arc, LazyLock,
Arc, LazyLock, RwLock as StdRwLock,
},
time::{Duration, Instant},
};
@@ -260,6 +260,18 @@ pub trait Worker: Send + Sync + fmt::Debug {
&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
/// Returns None for HTTP workers, Some(client) for gRPC workers
async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<GrpcClient>>>;
@@ -485,6 +497,10 @@ pub struct BasicWorker {
pub circuit_breaker: CircuitBreaker,
/// Lazily initialized gRPC client for gRPC workers
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 {
@@ -622,6 +638,40 @@ impl Worker for BasicWorker {
&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>>> {
match self.metadata.connection_mode {
ConnectionMode::Http => Ok(None),
+2 -1
View File
@@ -128,7 +128,7 @@ impl BasicWorkerBuilder {
pub fn build(self) -> BasicWorker {
use std::sync::{
atomic::{AtomicBool, AtomicUsize},
Arc,
Arc, RwLock as StdRwLock,
};
use tokio::sync::RwLock;
@@ -187,6 +187,7 @@ impl BasicWorkerBuilder {
consecutive_successes: Arc::new(AtomicUsize::new(0)),
circuit_breaker: CircuitBreaker::with_config(self.circuit_breaker_config),
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 uuid::Uuid;
use crate::core::{ConnectionMode, Worker, WorkerType};
use crate::core::{ConnectionMode, RuntimeType, Worker, WorkerType};
/// Unique identifier for a worker
#[derive(Debug, Clone, Hash, Eq, PartialEq)]
@@ -283,12 +283,14 @@ impl WorkerRegistry {
/// - model_id: Filter by specific model
/// - worker_type: Filter by worker type (Regular, Prefill, Decode)
/// - connection_mode: Filter by connection mode (Http, Grpc)
/// - runtime_type: Filter by runtime type (Sglang, Vllm, External)
/// - healthy_only: Only return healthy workers
pub fn get_workers_filtered(
&self,
model_id: Option<&str>,
worker_type: Option<WorkerType>,
connection_mode: Option<ConnectionMode>,
runtime_type: Option<RuntimeType>,
healthy_only: bool,
) -> Vec<Arc<dyn Worker>> {
// 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
if healthy_only && !w.is_healthy() {
return false;
@@ -222,6 +222,17 @@ impl StepExecutor for DiscoverModelsStep {
.get("worker_config")
.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);
let model_cards = fetch_models(&config.url, config.api_key.as_deref())
@@ -315,12 +326,6 @@ impl StepExecutor for CreateExternalWorkersStep {
.get("model_cards")
.ok_or_else(|| WorkflowError::ContextValueNotFound("model_cards".to_string()))?;
debug!(
"Creating {} external workers for {}",
model_cards.len(),
config.url
);
// Build configs from router settings
let circuit_breaker_config = {
let cfg = app_context.router_config.effective_circuit_breaker_config();
@@ -355,11 +360,14 @@ impl StepExecutor for CreateExternalWorkersStep {
// Normalize URL (ensure https:// for external APIs)
let normalized_url = normalize_external_url(&config.url);
// Create a worker for each model
let mut workers = Vec::new();
for model_card in model_cards.iter() {
// 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())
.model(model_card.clone())
.models(vec![]) // Empty models = accepts any model
.worker_type(WorkerType::Regular)
.connection_mode(ConnectionMode::Http)
.runtime_type(RuntimeType::External)
@@ -377,19 +385,54 @@ impl StepExecutor for CreateExternalWorkersStep {
let worker = Arc::new(builder.build()) as Arc<dyn Worker>;
worker.set_healthy(false);
debug!(
"Created external worker for model {} at {}",
model_card.id, normalized_url
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
);
info!(
"Created {} external workers from {}",
workers.len(),
config.url
);
// Create a worker for each model
for model_card in model_cards.iter() {
let mut builder = BasicWorkerBuilder::new(normalized_url.clone())
.model(model_card.clone())
.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);
debug!(
"Created external worker for model {} at {}",
model_card.id, normalized_url
);
workers.push(worker);
}
info!(
"Created {} external workers from {}",
workers.len(),
config.url
);
}
context.set("workers", workers);
context.set("labels", labels);
+7 -13
View File
@@ -56,9 +56,7 @@ impl RouterFactory {
)
.await
}
RoutingMode::OpenAI { worker_urls } => {
Self::create_openai_router(worker_urls.clone(), ctx).await
}
RoutingMode::OpenAI { .. } => Self::create_openai_router(ctx).await,
},
}
}
@@ -119,16 +117,12 @@ impl RouterFactory {
}
/// Create an OpenAI router
async fn create_openai_router(
worker_urls: Vec<String>,
ctx: &Arc<AppContext>,
) -> Result<Box<dyn RouterTrait>, String> {
if worker_urls.is_empty() {
return Err("OpenAI mode requires at least one worker URL".to_string());
}
let router = OpenAIRouter::new(worker_urls, ctx).await?;
///
/// Workers should be registered via the external worker registration workflow
/// before using this router. The workflow discovers models from the provided
/// endpoints and creates external workers in the registry.
async fn create_openai_router(ctx: &Arc<AppContext>) -> Result<Box<dyn RouterTrait>, String> {
let router = OpenAIRouter::new(ctx).await?;
Ok(Box::new(router))
}
}
@@ -120,6 +120,7 @@ impl WorkerSelectionStage {
model_id,
Some(WorkerType::Regular),
Some(ConnectionMode::Grpc { port: None }),
None, // any runtime type
false, // get all workers, we'll filter by is_available() next
);
@@ -153,6 +154,7 @@ impl WorkerSelectionStage {
model_id,
None,
Some(ConnectionMode::Grpc { port: None }), // Match any gRPC worker
None, // any runtime type
false,
);
+2
View File
@@ -137,12 +137,14 @@ impl std::fmt::Debug for GrpcPDRouter {
bootstrap_port: None,
}),
Some(ConnectionMode::Grpc { port: None }),
None,
false,
);
let decode_workers = self.worker_registry.get_workers_filtered(
None,
Some(WorkerType::Decode),
Some(ConnectionMode::Grpc { port: None }),
None,
false,
);
f.debug_struct("GrpcPDRouter")
+2
View File
@@ -53,6 +53,7 @@ impl Router {
None, // any model
Some(WorkerType::Regular),
Some(ConnectionMode::Http),
None, // any runtime type
false, // include all workers
);
@@ -139,6 +140,7 @@ impl Router {
effective_model_id,
Some(WorkerType::Regular),
Some(ConnectionMode::Http),
None, // any runtime type
false, // get all workers, we'll filter by is_available() next
);
+316 -241
View File
@@ -4,7 +4,6 @@ use std::{
any::Any,
collections::HashSet,
sync::{atomic::AtomicBool, Arc},
time::{Duration, Instant},
};
use axum::{
@@ -14,8 +13,7 @@ use axum::{
response::{IntoResponse, Response},
Json,
};
use dashmap::DashMap;
use futures_util::StreamExt;
use futures_util::{future::join_all, StreamExt};
use once_cell::sync::Lazy;
use serde_json::{json, to_value, Value};
use tokio::sync::mpsc;
@@ -35,10 +33,11 @@ use super::{
},
responses::{mask_tools_as_mcp, patch_streaming_response_json},
streaming::handle_streaming_response,
utils::{apply_provider_headers, extract_auth_header, probe_endpoint_for_model},
utils::{apply_provider_headers, extract_auth_header},
};
use crate::{
core::{CircuitBreaker, CircuitBreakerConfig as CoreCircuitBreakerConfig},
app_context::AppContext,
core::{ModelCard, RuntimeType, Worker, WorkerRegistry},
data_connector::{
ConversationId, ConversationItemStorage, ConversationStorage, ListParams, ResponseId,
ResponseStorage, SortOrder,
@@ -56,7 +55,6 @@ use crate::{
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
///
/// This router manages connections to OpenAI-compatible API endpoints (OpenAI, xAI, etc.)
/// using the Worker abstraction. Workers are registered via the external worker registration
/// workflow and stored in the WorkerRegistry.
pub struct OpenAIRouter {
/// HTTP client for upstream OpenAI-compatible API
client: reqwest::Client,
/// Multiple OpenAI-compatible API endpoints (OpenAI, xAI, etc.)
worker_urls: Vec<String>,
/// Model cache: model_id -> endpoint URL
model_cache: Arc<DashMap<String, CachedEndpoint>>,
/// Circuit breaker
circuit_breaker: CircuitBreaker,
/// Worker registry for model-based worker lookup
worker_registry: Arc<WorkerRegistry>,
/// Health status
healthy: AtomicBool,
/// Response storage for managing conversation history
@@ -120,8 +111,11 @@ pub struct OpenAIRouter {
impl std::fmt::Debug for OpenAIRouter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let registry_stats = self.worker_registry.stats();
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)
.finish()
}
@@ -131,33 +125,16 @@ impl OpenAIRouter {
/// Maximum number of conversation items to attach as input when a conversation is provided
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
pub async fn new(
worker_urls: Vec<String>,
ctx: &Arc<crate::app_context::AppContext>,
) -> Result<Self, String> {
///
/// Workers are registered separately via the external worker registration workflow.
/// This router queries the WorkerRegistry to find workers that support requested models.
pub async fn new(ctx: &Arc<AppContext>) -> Result<Self, String> {
// Use HTTP client from AppContext
let client = ctx.client.clone();
// Normalize URLs (remove trailing slashes)
let worker_urls: Vec<String> = worker_urls
.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 worker registry from AppContext
let worker_registry = ctx.worker_registry.clone();
// Get MCP manager from AppContext (must be initialized)
let mcp_manager = ctx
@@ -168,9 +145,7 @@ impl OpenAIRouter {
Ok(Self {
client,
worker_urls,
model_cache: Arc::new(DashMap::new()),
circuit_breaker,
worker_registry,
healthy: AtomicBool::new(true),
response_storage: ctx.response_storage.clone(),
conversation_storage: ctx.conversation_storage.clone(),
@@ -179,76 +154,178 @@ impl OpenAIRouter {
})
}
/// Discover which endpoint has the model
async fn find_endpoint_for_model(
/// Refresh models for a single external worker by querying its /v1/models endpoint.
///
/// Returns true if refresh succeeded and models were cached on the worker.
async fn refresh_worker_models(
&self,
worker: &Arc<dyn Worker>,
auth_header: Option<&HeaderValue>,
) -> bool {
let url = format!("{}/v1/models", worker.url());
// Build request to backend
let mut backend_req = self.client.get(&url);
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,
model_id: &str,
auth_header: Option<&str>,
) -> Result<String, Response> {
// Single endpoint - fast path
if self.worker_urls.len() == 1 {
return Ok(self.worker_urls[0].clone());
auth_header: Option<&HeaderValue>,
) -> Result<Arc<dyn Worker>, Box<Response>> {
// Helper to find candidates for a model
// Note: We get ALL external workers and filter by supports_model() because
// wildcard workers (empty models) aren't in the model index but support any model
let find_candidates = || {
self.worker_registry
.get_workers_filtered(
None, // Get all external workers, not just those in model index
None,
None,
Some(RuntimeType::External),
true, // healthy_only
)
.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
if let Some(entry) = self.model_cache.get(model_id) {
if entry.cached_at.elapsed() < Duration::from_secs(Self::MODEL_CACHE_TTL_SECS) {
return Ok(entry.url.clone());
}
// No match found - refresh models on all external workers
tracing::debug!(
"No worker found for model '{}', refreshing external worker models",
model_id
);
self.refresh_external_models(auth_header).await;
// Second try: check if any worker now supports the model after refresh
let candidates = find_candidates();
if !candidates.is_empty() {
return Ok(candidates
.into_iter()
.min_by_key(|w| w.load())
.expect("candidates is not empty"));
}
// 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);
}
}
// Model not found on any endpoint
Err((
StatusCode::NOT_FOUND,
Json(json!({
"error": {
"message": format!("Model '{}' not found on any endpoint", model_id),
"type": "model_not_found",
}
})),
)
.into_response())
Err(Box::new(
(
StatusCode::NOT_FOUND,
Json(json!({
"error": {
"message": format!("No worker available for model '{}'", model_id),
"type": "model_not_found",
}
})),
)
.into_response(),
))
}
/// Handle non-streaming response with optional MCP tool loop
async fn handle_non_streaming_response(
&self,
url: String,
worker: &Arc<dyn Worker>,
headers: Option<&HeaderMap>,
mut payload: Value,
original_body: &ResponsesRequest,
original_previous_response_id: Option<String>,
) -> Response {
let url = format!("{}/v1/responses", worker.url());
// Check if MCP is active for this request
// Ensure dynamic client is created if needed
if let Some(ref tools) = original_body.tools {
@@ -284,7 +361,7 @@ impl OpenAIRouter {
{
Ok(resp) => response_json = resp,
Err(err) => {
self.circuit_breaker.record_failure();
worker.circuit_breaker().record_failure();
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({"error": {"message": err}})),
@@ -296,14 +373,16 @@ impl OpenAIRouter {
// No MCP - simple request
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 {
Ok(r) => r,
Err(e) => {
self.circuit_breaker.record_failure();
worker.circuit_breaker().record_failure();
tracing::error!(
url = %url,
error = %e,
@@ -318,7 +397,7 @@ impl OpenAIRouter {
};
if !response.status().is_success() {
self.circuit_breaker.record_failure();
worker.circuit_breaker().record_failure();
let status = StatusCode::from_u16(response.status().as_u16())
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
let body = response.text().await.unwrap_or_default();
@@ -328,7 +407,7 @@ impl OpenAIRouter {
response_json = match response.json::<Value>().await {
Ok(r) => r,
Err(e) => {
self.circuit_breaker.record_failure();
worker.circuit_breaker().record_failure();
return (
StatusCode::INTERNAL_SERVER_ERROR,
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
@@ -376,134 +455,134 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
async fn health_generate(&self, _req: Request<Body>) -> Response {
// Check all endpoints in parallel - only healthy if ALL are healthy
if self.worker_urls.is_empty() {
return (StatusCode::SERVICE_UNAVAILABLE, "No endpoints configured").into_response();
// Check health of all external workers
let external_workers: Vec<_> = self
.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![];
for url in &self.worker_urls {
let url = url.clone();
let client = self.client.clone();
let mut healthy_count = 0;
let mut unhealthy_workers = Vec::new();
let handle = tokio::spawn(async move {
let probe_url = format!("{}/v1/models", url);
match client
.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 {
Err(format!("Endpoint {} returned status {}", url, code))
}
}
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)),
for worker in &external_workers {
if worker.is_healthy() {
healthy_count += 1;
} else {
unhealthy_workers.push(format!("{} ({})", worker.model_id(), worker.url()));
}
}
if errors.is_empty() {
(StatusCode::OK, "OK").into_response()
if unhealthy_workers.is_empty() {
(
StatusCode::OK,
format!("OK - {} workers healthy", healthy_count),
)
.into_response()
} else {
(
StatusCode::SERVICE_UNAVAILABLE,
format!("Some endpoints unhealthy: {}", errors.join(", ")),
format!(
"{}/{} workers unhealthy: {}",
unhealthy_workers.len(),
external_workers.len(),
unhealthy_workers.join(", ")
),
)
.into_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!({
"router_type": "openai",
"workers": self.worker_urls.len(),
"worker_urls": &self.worker_urls
"total_workers": stats.total_workers,
"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()
}
async fn get_models(&self, req: Request<Body>) -> Response {
// Aggregate models from all endpoints
if self.worker_urls.is_empty() {
return (StatusCode::SERVICE_UNAVAILABLE, "No endpoints configured").into_response();
// Return models from all registered external workers
let external_workers: Vec<_> = self
.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();
let auth = headers
.get("authorization")
.or_else(|| headers.get("Authorization"));
// Refresh models for all external workers using user's auth header
let auth_header = extract_auth_header(Some(req.headers()), &None);
self.refresh_external_models(auth_header.as_ref()).await;
// Query all endpoints in parallel
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
// Collect models from all workers
let mut all_models = Vec::new();
for handle in handles {
if let Ok(Ok(json)) = handle.await {
if let Some(data) = json.get("data").and_then(|v| v.as_array()) {
all_models.extend_from_slice(data);
let mut seen_models = HashSet::new();
for worker in &external_workers {
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,
_model_id: Option<&str>,
) -> Response {
if !self.circuit_breaker.can_execute() {
return (StatusCode::SERVICE_UNAVAILABLE, "Circuit breaker open").into_response();
}
// Extract auth header for passthrough mode
let auth_header = extract_auth_header(headers, &None);
// Extract auth header
let auth = extract_auth_header(headers);
// Find endpoint for model
let base_url = match self
.find_endpoint_for_model(body.model.as_str(), auth)
// Select worker for model (discovery happens inside if needed)
let worker = match self
.select_worker_for_model(body.model.as_str(), auth_header.as_ref())
.await
{
Ok(url) => url,
Err(response) => return response,
Ok(w) => w,
Err(response) => return *response,
};
// 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);
// Forward Authorization header if provided
if let Some(h) = headers {
if let Some(auth) = h.get("authorization").or_else(|| h.get("Authorization")) {
req = req.header("Authorization", auth);
}
}
// 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());
req = apply_provider_headers(req, &url, auth_header.as_ref());
// Accept SSE when stream=true
if body.stream {
@@ -600,7 +673,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
let resp = match req.send().await {
Ok(r) => r,
Err(e) => {
self.circuit_breaker.record_failure();
worker.circuit_breaker().record_failure();
return (
StatusCode::SERVICE_UNAVAILABLE,
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();
match resp.bytes().await {
Ok(body) => {
self.circuit_breaker.record_success();
worker.circuit_breaker().record_success();
let mut response = Response::new(Body::from(body));
*response.status_mut() = status;
if let Some(ct) = content_type {
@@ -626,7 +699,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
response
}
Err(e) => {
self.circuit_breaker.record_failure();
worker.circuit_breaker().record_failure();
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("Failed to read response: {}", e),
@@ -683,18 +756,19 @@ impl crate::routers::RouterTrait for OpenAIRouter {
body: &ResponsesRequest,
model_id: Option<&str>,
) -> Response {
// Extract auth header
let auth = extract_auth_header(headers);
// Extract auth header for passthrough mode
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 base_url = match self.find_endpoint_for_model(model, auth).await {
Ok(url) => url,
Err(response) => return response,
let worker = match self
.select_worker_for_model(model, auth_header.as_ref())
.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
let mut request_body = body.clone();
if let Some(model) = model_id {
@@ -992,10 +1066,11 @@ impl crate::routers::RouterTrait for OpenAIRouter {
}
// Delegate to streaming or non-streaming handler
let url = format!("{}/v1/responses", worker.url());
if body.stream.unwrap_or(false) {
handle_streaming_response(
&self.client,
&self.circuit_breaker,
worker.circuit_breaker(),
Some(&self.mcp_manager),
self.response_storage.clone(),
self.conversation_storage.clone(),
@@ -1009,7 +1084,7 @@ impl crate::routers::RouterTrait for OpenAIRouter {
.await
} else {
self.handle_non_streaming_response(
url,
&worker,
headers,
payload,
body,
+28 -59
View File
@@ -2,7 +2,7 @@
use std::collections::HashMap;
use axum::http::{HeaderMap, HeaderValue};
use axum::http::HeaderValue;
// ============================================================================
// SSE Event Type Constants
@@ -99,16 +99,6 @@ impl OutputIndexMapper {
// 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
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ApiProvider {
@@ -168,56 +158,35 @@ pub fn apply_provider_headers(
req
}
/// Probe a single endpoint to check if it has the model
/// Returns Ok(url) if model found, Err(()) otherwise
pub async fn probe_endpoint_for_model(
client: reqwest::Client,
url: String,
model: String,
auth: Option<String>,
) -> Result<String, ()> {
use tracing::debug;
// ============================================================================
// Auth Header Resolution
// ============================================================================
let probe_url = format!("{}/v1/models/{}", url, model);
let req = client
.get(&probe_url)
.timeout(std::time::Duration::from_secs(5));
/// Extract auth header with passthrough semantics.
///
/// Passthrough mode: User's Authorization header takes priority.
/// 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.)
let auth_header_value = auth.as_ref().and_then(|a| HeaderValue::from_str(a).ok());
let req = apply_provider_headers(req, &url, auth_header_value.as_ref());
match req.send().await {
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(())
}
}
// Return user's auth if provided, otherwise use worker's API key
user_auth.or_else(|| {
worker_api_key
.as_ref()
.and_then(|k| HeaderValue::from_str(&format!("Bearer {}", k)).ok())
})
}
// ============================================================================
+44 -2
View File
@@ -16,8 +16,10 @@ use std::{
use serde_json::json;
use sgl_model_gateway::{
app_context::AppContext,
config::RouterConfig,
core::{LoadMonitor, WorkerRegistry},
config::{RouterConfig, RoutingMode},
core::{
BasicWorkerBuilder, LoadMonitor, ModelCard, RuntimeType, Worker, WorkerRegistry, WorkerType,
},
data_connector::{
MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage,
},
@@ -111,6 +113,26 @@ pub async fn create_test_context(config: RouterConfig) -> Arc<AppContext> {
.set(engine)
.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
use sgl_model_gateway::mcp::{McpConfig, McpManager};
let empty_config = McpConfig {
@@ -222,6 +244,26 @@ pub async fn create_test_context_with_mcp_config(
.set(engine)
.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
let mcp_config = McpConfig::from_file(mcp_config_path)
.await
+50 -1
View File
@@ -5,7 +5,9 @@ use reqwest::Client;
use sgl_model_gateway::{
app_context::AppContext,
config::RouterConfig,
core::{LoadMonitor, WorkerRegistry},
core::{
BasicWorkerBuilder, LoadMonitor, ModelCard, RuntimeType, Worker, WorkerRegistry, WorkerType,
},
data_connector::{
MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage,
},
@@ -209,3 +211,50 @@ pub async fn create_test_app_context() -> Arc<AppContext> {
.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]
async fn test_openai_router_creation() {
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");
@@ -109,9 +111,8 @@ async fn test_openai_router_creation() {
#[tokio::test]
async fn test_openai_router_server_info() {
let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec!["https://api.openai.com".to_string()], &ctx)
.await
.unwrap();
common::test_app::register_external_worker(&ctx, "https://api.openai.com", None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
let req = Request::builder()
.method(Method::GET)
@@ -135,9 +136,8 @@ async fn test_openai_router_models() {
// Use mock server for deterministic models response
let mock_server = MockOpenAIServer::new().await;
let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec![mock_server.base_url()], &ctx)
.await
.unwrap();
common::test_app::register_external_worker(&ctx, &mock_server.base_url(), None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
let req = Request::builder()
.method(Method::GET)
@@ -209,7 +209,8 @@ async fn test_openai_router_responses_with_mock() {
let base_url = format!("http://{}", addr);
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)
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 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
let storage = ctx.response_storage.clone();
@@ -598,9 +600,8 @@ async fn test_router_factory_openai_mode() {
#[tokio::test]
async fn test_unsupported_endpoints() {
let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec!["https://api.openai.com".to_string()], &ctx)
.await
.unwrap();
common::test_app::register_external_worker(&ctx, "https://api.openai.com", None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
let generate_request = GenerateRequest {
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 ctx = common::test_app::create_test_app_context().await;
// Create router pointing to mock server
let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap();
// Register the mock server worker and create router
common::test_app::register_external_worker(&ctx, &base_url, None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
// Create a minimal chat completion 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 ctx = common::test_app::create_test_app_context().await;
// Create router
let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap();
// Register the mock server worker and create router
common::test_app::register_external_worker(&ctx, &base_url, None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
// Create Axum app with chat completions endpoint
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 base_url = mock_server.base_url();
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
let val = json!({
@@ -797,9 +801,8 @@ async fn test_openai_router_chat_streaming_with_mock() {
#[tokio::test]
async fn test_openai_router_circuit_breaker() {
let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec!["http://invalid-url-that-will-fail".to_string()], &ctx)
.await
.unwrap();
common::test_app::register_external_worker(&ctx, "http://invalid-url-that-will-fail", None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
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]
async fn test_openai_router_models_auth_forwarding() {
// 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;
async fn test_openai_router_models_from_registry() {
let ctx = common::test_app::create_test_app_context().await;
let router = OpenAIRouter::new(vec![mock_server.base_url()], &ctx)
.await
.unwrap();
// Register a worker with the default model
common::test_app::register_external_worker(&ctx, "https://api.example.com", None);
let router = OpenAIRouter::new(&ctx).await.unwrap();
// 1) Without auth header -> expect 200 with empty model list
// (multi-endpoint aggregation silently skips failed endpoints)
// Get models - should return the registered model
let req = Request::builder()
.method(Method::GET)
.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 models: serde_json::Value = serde_json::from_str(&body_str).unwrap();
assert_eq!(models["object"], "list");
assert_eq!(models["data"].as_array().unwrap().len(), 0); // Empty when auth fails
// 2) With auth header -> expect 200
let req = Request::builder()
.method(Method::GET)
.uri("/models")
.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");
// Should have the default model (gpt-3.5-turbo)
let data = models["data"].as_array().unwrap();
assert_eq!(data.len(), 1);
assert_eq!(data[0]["id"], "gpt-3.5-turbo");
}
#[test]