Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion bindings/golang/src/policy.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,10 @@ use std::{

use async_trait::async_trait;
use llm_tokenizer::{create_tokenizer_from_file, traits::Tokenizer};
use openai_protocol::{chat::ChatCompletionRequest, worker::WorkerSpec};
use openai_protocol::{
chat::ChatCompletionRequest,
worker::{HealthCheckConfig, WorkerSpec},
};
use smg::{
core::{
circuit_breaker::CircuitBreaker,
Expand Down Expand Up @@ -63,6 +66,7 @@ impl GrpcWorker {

let metadata = WorkerMetadata {
spec,
health_config: HealthCheckConfig::default(),
health_endpoint: "/health".to_string(),
default_model_type: ModelType::LLM,
};
Expand Down
12 changes: 8 additions & 4 deletions model_gateway/src/core/job_queue.rs
Original file line number Diff line number Diff line change
Expand Up @@ -612,9 +612,11 @@ impl JobQueue {
spec.worker_type = proto_worker_type;
spec.api_key = api_key.clone();
spec.bootstrap_port = bootstrap_port;
spec.health = router_config.health_check.to_protocol_config();
// Health config is resolved at worker build time from router
// defaults + per-worker overrides (spec.health). No need to
// set spec.health here since these workers have no overrides.
spec.max_connection_attempts =
router_config.health_check.success_threshold * 10;
router_config.health_check.success_threshold.max(1) * 10;
let config = spec;

let job = Job::AddWorker {
Expand Down Expand Up @@ -804,7 +806,9 @@ fn build_external_worker_config(
let mut spec = WorkerSpec::new(url);
spec.runtime_type = RuntimeType::External;
spec.api_key = api_key;
spec.health = router_config.health_check.to_protocol_config();
spec.max_connection_attempts = router_config.health_check.success_threshold * 10;
// Health config is resolved at worker build time from router
// defaults + per-worker overrides (spec.health). No need to
// set spec.health here since these workers have no overrides.
spec.max_connection_attempts = router_config.health_check.success_threshold.max(1) * 10;
spec
}
23 changes: 12 additions & 11 deletions model_gateway/src/core/steps/worker/external/create_workers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
use std::{collections::HashMap, sync::Arc, time::Duration};

use async_trait::async_trait;
use openai_protocol::worker::HealthCheckConfig;
use tracing::{debug, info};
use wfaas::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult};

Expand Down Expand Up @@ -52,16 +51,18 @@ impl StepExecutor<ExternalWorkerWorkflowData> for CreateExternalWorkersStep {
};

let (health_config, health_endpoint) = {
let cfg = &app_context.router_config.health_check;
let protocol_config = HealthCheckConfig {
timeout_secs: cfg.timeout_secs,
check_interval_secs: cfg.check_interval_secs,
success_threshold: cfg.success_threshold,
failure_threshold: cfg.failure_threshold,
disable_health_check: cfg.disable_health_check
|| config.health.disable_health_check,
};
(protocol_config, cfg.endpoint.clone())
let base = app_context.router_config.health_check.to_protocol_config();
let mut merged = config.health.apply_to(&base);
// External workers (OpenAI, Anthropic, etc.) should not be health-checked
// by default — they are third-party APIs that don't expose a /health endpoint.
// Only apply the default if the user hasn't explicitly overridden it.
if config.health.disable_health_check.is_none() {
merged.disable_health_check = true;
}
(
merged,
app_context.router_config.health_check.endpoint.clone(),
)
};
Comment thread
coderabbitai[bot] marked this conversation as resolved.

// Build labels from config
Expand Down
43 changes: 14 additions & 29 deletions model_gateway/src/core/steps/worker/local/create_worker.rs
Original file line number Diff line number Diff line change
Expand Up @@ -171,26 +171,14 @@ fn build_model_card(
config: &WorkerSpec,
labels: &HashMap<String, String>,
) -> ModelCard {
let mut card = ModelCard::new(model_id);

// If the user provided a matching model in config.models, copy its fields
if let Some(proto_card) = config.models.find(model_id) {
if let Some(ref tokenizer_path) = proto_card.tokenizer_path {
card = card.with_tokenizer_path(tokenizer_path.clone());
}
if let Some(ref reasoning_parser) = proto_card.reasoning_parser {
card = card.with_reasoning_parser(reasoning_parser.clone());
}
if let Some(ref tool_parser) = proto_card.tool_parser {
card = card.with_tool_parser(tool_parser.clone());
}
if let Some(ref chat_template) = proto_card.chat_template {
card = card.with_chat_template(chat_template.clone());
}
if !proto_card.aliases.is_empty() {
card = card.with_aliases(proto_card.aliases.clone());
}
}
// Start from the user-provided proto_card if it matches, preserving all
// user-supplied fields (display_name, provider, context_length, etc.).
// Otherwise start from a blank card with just the model_id.
let mut card = config
.models
.find(model_id)
.cloned()
.unwrap_or_else(|| ModelCard::new(model_id));
if let Some(model_type_str) = labels.get("model_type") {
card = card.with_hf_model_type(model_type_str.clone());
}
Expand Down Expand Up @@ -277,15 +265,12 @@ fn build_health_config(
app_context: &AppContext,
config: &WorkerSpec,
) -> (HealthCheckConfig, String) {
let cfg = &app_context.router_config.health_check;
let protocol_config = HealthCheckConfig {
timeout_secs: cfg.timeout_secs,
check_interval_secs: cfg.check_interval_secs,
success_threshold: cfg.success_threshold,
failure_threshold: cfg.failure_threshold,
disable_health_check: cfg.disable_health_check || config.health.disable_health_check,
};
(protocol_config, cfg.endpoint.clone())
let base = app_context.router_config.health_check.to_protocol_config();
let merged = config.health.apply_to(&base);
(
merged,
app_context.router_config.health_check.endpoint.clone(),
)
}

fn normalize_url(url: &str, connection_mode: &ConnectionMode) -> String {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -115,13 +115,17 @@ impl StepExecutor<LocalWorkerWorkflowData> for DetectConnectionModeStep {
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;

debug!(
"Detecting connection mode for {} (timeout: {}s, max_attempts: {})",
"Detecting connection mode for {} (timeout: {:?}s, max_attempts: {})",
config.url, config.health.timeout_secs, config.max_connection_attempts
);

// Try both protocols in parallel
let url = config.url.clone();
let timeout = config.health.timeout_secs;
// Use per-worker timeout override if set, otherwise fall back to router default
let timeout = config
.health
.timeout_secs
.unwrap_or(app_context.router_config.health_check.timeout_secs);
let client = &app_context.client;
Comment thread
slin1237 marked this conversation as resolved.
// Auto-detect runtime unless explicitly set to non-default
let runtime_type_str = config.runtime_type.to_string();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -54,8 +54,8 @@ impl StepExecutor<WorkerUpdateWorkflowData> for UpdateWorkerPropertiesStep {
let updated_priority = request.priority.unwrap_or(worker.priority());
let updated_cost = request.cost.unwrap_or(worker.cost());

// Build updated health config
let existing_health = &worker.metadata().spec.health;
// Build updated health config from resolved runtime config
let existing_health = &worker.metadata().health_config;
let updated_health_config = match &request.health {
Some(update) => update.apply_to(existing_health),
None => existing_health.clone(),
Expand Down
50 changes: 32 additions & 18 deletions model_gateway/src/core/worker.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@ pub use openai_protocol::worker::{ConnectionMode, RuntimeType, WorkerType};
use openai_protocol::{
model_card::ModelCard,
model_type::{Endpoint, ModelType},
worker::{ProviderType, WorkerInfo, WorkerModels, WorkerSpec},
worker::{HealthCheckConfig, ProviderType, WorkerInfo, WorkerModels, WorkerSpec},
};
use tokio::{sync::OnceCell, time};

Expand Down Expand Up @@ -401,6 +401,10 @@ impl WorkerTypeExt for WorkerType {
pub struct WorkerMetadata {
/// Protocol-level worker identity and configuration.
pub spec: WorkerSpec,
/// Resolved health check config (router defaults + per-worker overrides).
/// This is the concrete config used at runtime; `spec.health` only stores
/// the partial overrides from the API layer.
pub health_config: HealthCheckConfig,
/// Health check endpoint path (internal-only, from router config).
pub health_endpoint: String,
/// Default model type for unknown models (defaults to LLM capabilities).
Expand Down Expand Up @@ -529,7 +533,7 @@ impl Worker for BasicWorker {
}

async fn check_health_async(&self) -> WorkerResult<()> {
if self.metadata.spec.health.disable_health_check {
if self.metadata.health_config.disable_health_check {
if !self.is_healthy() {
self.set_healthy(true);
}
Expand All @@ -552,7 +556,7 @@ impl Worker for BasicWorker {
Metrics::record_worker_health_check(worker_type_str, metrics_labels::CB_SUCCESS);

if !self.is_healthy()
&& successes >= self.metadata.spec.health.success_threshold as usize
&& successes >= self.metadata.health_config.success_threshold as usize
{
self.set_healthy(true);
self.consecutive_successes.store(0, Ordering::Release);
Expand All @@ -565,7 +569,8 @@ impl Worker for BasicWorker {
// Record health check failure metric
Metrics::record_worker_health_check(worker_type_str, metrics_labels::CB_FAILURE);

if self.is_healthy() && failures >= self.metadata.spec.health.failure_threshold as usize
if self.is_healthy()
&& failures >= self.metadata.health_config.failure_threshold as usize
{
self.set_healthy(false);
self.consecutive_failures.store(0, Ordering::Release);
Expand Down Expand Up @@ -751,7 +756,7 @@ impl Worker for BasicWorker {
}

async fn grpc_health_check(&self) -> WorkerResult<bool> {
let timeout = Duration::from_secs(self.metadata.spec.health.timeout_secs);
let timeout = Duration::from_secs(self.metadata.health_config.timeout_secs);
let maybe = self.get_grpc_client().await?;
let Some(grpc_client) = maybe else {
tracing::error!(
Expand Down Expand Up @@ -785,7 +790,7 @@ impl Worker for BasicWorker {
}

async fn http_health_check(&self) -> WorkerResult<bool> {
let timeout = Duration::from_secs(self.metadata.spec.health.timeout_secs);
let timeout = Duration::from_secs(self.metadata.health_config.timeout_secs);

let health_url = format!("{}{}", self.base_url(), self.metadata.health_endpoint);

Expand Down Expand Up @@ -905,31 +910,38 @@ impl<T: Send + Unpin + 'static> http_body::Body for AttachedBody<T> {
}
}

/// Health checker handle with graceful shutdown
/// Health checker handle with graceful shutdown.
///
/// The checker sleeps until the next worker is due for a health check,
/// so it wakes only when there is actual work to do.
pub(crate) struct HealthChecker {
#[allow(dead_code)]
handle: tokio::task::JoinHandle<()>,
shutdown: Arc<AtomicBool>,
shutdown_notify: Arc<tokio::sync::Notify>,
}

impl fmt::Debug for HealthChecker {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("HealthChecker")
.field("shutdown", &self.shutdown.load(Ordering::Relaxed))
.finish()
f.debug_struct("HealthChecker").finish()
}
}

impl HealthChecker {
/// Create a new HealthChecker
pub fn new(handle: tokio::task::JoinHandle<()>, shutdown: Arc<AtomicBool>) -> Self {
Self { handle, shutdown }
pub fn new(
handle: tokio::task::JoinHandle<()>,
shutdown_notify: Arc<tokio::sync::Notify>,
) -> Self {
Self {
handle,
shutdown_notify,
}
}

/// Shutdown the health checker gracefully
/// Shutdown the health checker gracefully.
/// Wakes the sleeping task immediately so it can exit.
#[allow(dead_code)]
pub async fn shutdown(self) {
self.shutdown.store(true, Ordering::Release);
self.shutdown_notify.notify_one();
let _ = self.handle.await;
}
}
Expand Down Expand Up @@ -1053,8 +1065,8 @@ mod tests {
.health_endpoint("/custom-health")
.build();

assert_eq!(worker.metadata().spec.health.timeout_secs, 15);
assert_eq!(worker.metadata().spec.health.check_interval_secs, 45);
assert_eq!(worker.metadata().health_config.timeout_secs, 15);
assert_eq!(worker.metadata().health_config.check_interval_secs, 45);
assert_eq!(worker.metadata().health_endpoint, "/custom-health");
}

Expand Down Expand Up @@ -1573,6 +1585,7 @@ mod tests {
fn test_worker_metadata_empty_models_accepts_all() {
let metadata = WorkerMetadata {
spec: WorkerSpec::new("http://test:8080"),
health_config: HealthCheckConfig::default(),
health_endpoint: "/health".to_string(),
default_model_type: ModelType::LLM,
};
Expand All @@ -1596,6 +1609,7 @@ mod tests {
spec.models = WorkerModels::from(vec![model1, model2]);
let metadata = WorkerMetadata {
spec,
health_config: HealthCheckConfig::default(),
health_endpoint: "/health".to_string(),
default_model_type: ModelType::LLM,
};
Expand Down
Loading