diff --git a/sgl-router/src/core/worker.rs b/sgl-router/src/core/worker.rs index c1457df57..697d21a95 100644 --- a/sgl-router/src/core/worker.rs +++ b/sgl-router/src/core/worker.rs @@ -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) { + // 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>>; @@ -485,6 +497,10 @@ pub struct BasicWorker { pub circuit_breaker: CircuitBreaker, /// Lazily initialized gRPC client for gRPC workers pub grpc_client: Arc>>>, + /// 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>>>, } 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) { + 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>> { match self.metadata.connection_mode { ConnectionMode::Http => Ok(None), diff --git a/sgl-router/src/core/worker_builder.rs b/sgl-router/src/core/worker_builder.rs index de0354950..a650e2f95 100644 --- a/sgl-router/src/core/worker_builder.rs +++ b/sgl-router/src/core/worker_builder.rs @@ -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)), } } } diff --git a/sgl-router/src/core/worker_registry.rs b/sgl-router/src/core/worker_registry.rs index 95e09f87c..3648906e1 100644 --- a/sgl-router/src/core/worker_registry.rs +++ b/sgl-router/src/core/worker_registry.rs @@ -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, connection_mode: Option, + runtime_type: Option, healthy_only: bool, ) -> Vec> { // 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; diff --git a/sgl-router/src/core/workflow/steps/external_worker_registration.rs b/sgl-router/src/core/workflow/steps/external_worker_registration.rs index aaefac8c1..4ebbefc4d 100644 --- a/sgl-router/src/core/workflow/steps/external_worker_registration.rs +++ b/sgl-router/src/core/workflow/steps/external_worker_registration.rs @@ -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::>("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; 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; + 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); diff --git a/sgl-router/src/routers/factory.rs b/sgl-router/src/routers/factory.rs index 5f6c962f5..a4c85d5de 100644 --- a/sgl-router/src/routers/factory.rs +++ b/sgl-router/src/routers/factory.rs @@ -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, - ctx: &Arc, - ) -> Result, 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) -> Result, String> { + let router = OpenAIRouter::new(ctx).await?; Ok(Box::new(router)) } } diff --git a/sgl-router/src/routers/grpc/common/stages/worker_selection.rs b/sgl-router/src/routers/grpc/common/stages/worker_selection.rs index d52819374..969a45294 100644 --- a/sgl-router/src/routers/grpc/common/stages/worker_selection.rs +++ b/sgl-router/src/routers/grpc/common/stages/worker_selection.rs @@ -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, ); diff --git a/sgl-router/src/routers/grpc/pd_router.rs b/sgl-router/src/routers/grpc/pd_router.rs index 953804c01..c9c9afdc3 100644 --- a/sgl-router/src/routers/grpc/pd_router.rs +++ b/sgl-router/src/routers/grpc/pd_router.rs @@ -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") diff --git a/sgl-router/src/routers/http/router.rs b/sgl-router/src/routers/http/router.rs index e9f4dba40..ca9de0b92 100644 --- a/sgl-router/src/routers/http/router.rs +++ b/sgl-router/src/routers/http/router.rs @@ -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 ); diff --git a/sgl-router/src/routers/openai/router.rs b/sgl-router/src/routers/openai/router.rs index 01ec0b626..773648590 100644 --- a/sgl-router/src/routers/openai/router.rs +++ b/sgl-router/src/routers/openai/router.rs @@ -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> = 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, - /// Model cache: model_id -> endpoint URL - model_cache: Arc>, - /// Circuit breaker - circuit_breaker: CircuitBreaker, + /// Worker registry for model-based worker lookup + worker_registry: Arc, /// 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", ®istry_stats.total_workers) + .field("registered_models", ®istry_stats.total_models) + .field("healthy_workers", ®istry_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, - ctx: &Arc, - ) -> Result { + /// + /// 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) -> Result { // Use HTTP client from AppContext let client = ctx.client.clone(); - // Normalize URLs (remove trailing slashes) - let worker_urls: Vec = 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, + 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::().await { + Ok(json_response) => { + if let Some(data) = json_response.get("data").and_then(|d| d.as_array()) { + let model_cards: Vec = 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 { - // Single endpoint - fast path - if self.worker_urls.len() == 1 { - return Ok(self.worker_urls[0].clone()); + auth_header: Option<&HeaderValue>, + ) -> Result, Box> { + // 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::>() + }; + + // 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, headers: Option<&HeaderMap>, mut payload: Value, original_body: &ResponsesRequest, original_previous_response_id: Option, ) -> 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::().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) -> 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) -> 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 = 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) -> 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::().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, diff --git a/sgl-router/src/routers/openai/utils.rs b/sgl-router/src/routers/openai/utils.rs index aa1a80b25..f7f9ead8c 100644 --- a/sgl-router/src/routers/openai/utils.rs +++ b/sgl-router/src/routers/openai/utils.rs @@ -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, -) -> Result { - 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, +) -> Option { + // 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()) + }) } // ============================================================================ diff --git a/sgl-router/tests/common/mod.rs b/sgl-router/tests/common/mod.rs index aa3893fb2..1a8b01a9d 100644 --- a/sgl-router/tests/common/mod.rs +++ b/sgl-router/tests/common/mod.rs @@ -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 { .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 = 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 = 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 diff --git a/sgl-router/tests/common/test_app.rs b/sgl-router/tests/common/test_app.rs index 2fbb8c74b..74de6dae2 100644 --- a/sgl-router/tests/common/test_app.rs +++ b/sgl-router/tests/common/test_app.rs @@ -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 { .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, url: &str, models: Option>) { + let model_list: Vec = models + .unwrap_or_else(|| vec!["gpt-3.5-turbo"]) + .into_iter() + .map(ModelCard::new) + .collect(); + + let worker: Arc = 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, url: &str, model_card: ModelCard) { + let worker: Arc = Arc::new( + BasicWorkerBuilder::new(url) + .worker_type(WorkerType::Regular) + .runtime_type(RuntimeType::External) + .model(model_card) + .build(), + ); + + ctx.worker_registry.register(worker); +} diff --git a/sgl-router/tests/test_openai_routing.rs b/sgl-router/tests/test_openai_routing.rs index 6c24f9bf2..282522d5e 100644 --- a/sgl-router/tests/test_openai_routing.rs +++ b/sgl-router/tests/test_openai_routing.rs @@ -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]