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