[model-gateway] use worker crate in openai router (#14330)
This commit is contained in:
@@ -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),
|
||||
|
||||
@@ -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)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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,
|
||||
);
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
);
|
||||
|
||||
|
||||
@@ -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", ®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<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,
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user