[SMG] Add /v1/models fallback for model name discovery (#25293)

Co-authored-by: Amit Gruner <agruner@crusoe.ai>
This commit is contained in:
gruner
2026-05-18 22:02:35 +08:00
committed by GitHub
co-authored by Amit Gruner
parent ba2ffcf156
commit 0ab427d0e1
4 changed files with 233 additions and 0 deletions
@@ -1781,3 +1781,93 @@ impl Default for MockWorkerConfig {
}
}
}
/// A minimal OpenAI-compatible mock worker that does not implement /server_info or /model_info.
/// Used to test fallback model name discovery via /v1/models.
pub struct OpenAiOnlyMockWorker {
port: u16,
model_name: String,
shutdown_handle: Option<tokio::task::JoinHandle<()>>,
shutdown_tx: Option<tokio::sync::oneshot::Sender<()>>,
}
impl OpenAiOnlyMockWorker {
pub fn new(model_name: impl Into<String>) -> Self {
Self {
port: 0,
model_name: model_name.into(),
shutdown_handle: None,
shutdown_tx: None,
}
}
pub async fn start(&mut self) -> Result<String, Box<dyn std::error::Error>> {
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
self.port = listener.local_addr()?.port();
drop(listener);
let model_name = self.model_name.clone();
let port = self.port;
let app = Router::new()
.route("/health", get(|| async { Json(json!({ "status": "healthy" })) }))
.route("/health_generate", get(|| async { Json(json!({ "status": "ok" })) }))
.route(
"/v1/models",
get(move || {
let model_name = model_name.clone();
async move {
let ts = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
Json(json!({
"object": "list",
"data": [{ "id": model_name, "object": "model", "created": ts, "owned_by": "owner" }]
}))
}
}),
);
let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>();
self.shutdown_tx = Some(shutdown_tx);
let handle = tokio::spawn(async move {
let listener = match tokio::net::TcpListener::bind(("127.0.0.1", port)).await {
Ok(l) => l,
Err(e) => {
eprintln!("Failed to bind to port {}: {}", port, e);
return;
}
};
let server = axum::serve(listener, app).with_graceful_shutdown(async move {
let _ = shutdown_rx.await;
});
if let Err(e) = server.await {
eprintln!("Server error: {}", e);
}
});
self.shutdown_handle = Some(handle);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
Ok(format!("http://127.0.0.1:{}", self.port))
}
pub async fn stop(&mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if let Some(h) = self.shutdown_handle.take() {
let _ = tokio::time::timeout(tokio::time::Duration::from_secs(5), h).await;
}
}
}
impl Drop for OpenAiOnlyMockWorker {
fn drop(&mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
}
}
+1
View File
@@ -11,4 +11,5 @@ pub mod power_of_two_test;
pub mod service_discovery_test;
pub mod test_openai_routing;
pub mod test_pd_routing;
pub mod worker_discovery_test;
pub mod worker_management_test;
@@ -0,0 +1,94 @@
//! Worker metadata discovery integration tests.
use smg::{config::RouterConfig, core::Job};
use crate::common::{
create_test_context,
mock_worker::{HealthStatus, MockWorkerConfig, OpenAiOnlyMockWorker, WorkerType},
AppTestContext,
};
#[cfg(test)]
mod worker_discovery_tests {
use super::*;
/// Normal path: model name is discovered from /server_info.
#[tokio::test]
async fn test_model_name_discovered_via_server_info() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 0,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 0.0,
}])
.await;
let discovered_models = ctx.app_context.worker_registry.get_models();
assert!(
discovered_models.contains(&"mock-model-path".to_string()),
"Expected 'mock-model-path' discovered via /server_info, got: {:?}",
discovered_models
);
ctx.shutdown().await;
}
/// Fallback path: when /server_info is unavailable, model name is discovered via /v1/models.
#[tokio::test]
async fn test_model_name_discovered_via_v1_models_fallback() {
let mut worker = OpenAiOnlyMockWorker::new("my-model");
let url = worker.start().await.unwrap();
let config = RouterConfig::builder()
.regular_mode(vec![url.clone()])
.random_policy()
.host("127.0.0.1")
.port(0)
.max_payload_size(256 * 1024 * 1024)
.request_timeout_secs(600)
.worker_startup_timeout_secs(5)
.worker_startup_check_interval_secs(1)
.max_concurrent_requests(64)
.queue_timeout_secs(60)
.build_unchecked();
let app_context = create_test_context(config.clone()).await;
let job_queue = app_context
.worker_job_queue
.get()
.expect("JobQueue should be initialized");
job_queue
.submit(Job::InitializeWorkersFromConfig {
router_config: Box::new(config),
})
.await
.expect("Failed to submit worker initialization job");
let start = tokio::time::Instant::now();
loop {
if app_context
.worker_registry
.get_all()
.iter()
.any(|w| w.is_healthy())
{
break;
}
if start.elapsed().as_secs() > 10 {
panic!("Timeout waiting for worker to become healthy");
}
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
}
let discovered_models = app_context.worker_registry.get_models();
assert!(
discovered_models.contains(&"my-model".to_string()),
"Expected 'my-model' discovered via /v1/models fallback, got: {:?}",
discovered_models
);
worker.stop().await;
}
}