[SMG] Add /v1/models fallback for model name discovery (#25293)
Co-authored-by: Amit Gruner <agruner@crusoe.ai>
This commit is contained in:
@@ -169,6 +169,45 @@ pub async fn get_model_info(url: &str, api_key: Option<&str>) -> Result<ModelInf
|
||||
.map_err(|e| format!("Failed to parse response from {}: {}", model_info_url, e))
|
||||
}
|
||||
|
||||
/// Get model name from /v1/models endpoint (OpenAI-compatible fallback).
|
||||
async fn get_model_name_from_v1_models(url: &str, api_key: Option<&str>) -> Result<String, String> {
|
||||
let base_url = url.trim_end_matches('/');
|
||||
let models_url = format!("{}/v1/models", base_url);
|
||||
|
||||
let mut req = HTTP_CLIENT.get(&models_url);
|
||||
if let Some(key) = api_key {
|
||||
req = req.bearer_auth(key);
|
||||
}
|
||||
|
||||
let response = req
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("Failed to connect to {}: {}", models_url, e))?;
|
||||
|
||||
if !response.status().is_success() {
|
||||
return Err(format!(
|
||||
"Server returned status {} from {}",
|
||||
response.status(),
|
||||
models_url
|
||||
));
|
||||
}
|
||||
|
||||
let json: Value = response
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| format!("Failed to parse response from {}: {}", models_url, e))?;
|
||||
|
||||
json["data"]
|
||||
.as_array()
|
||||
.and_then(|arr| {
|
||||
arr.iter()
|
||||
.find(|entry| entry["object"].as_str() == Some("model"))
|
||||
})
|
||||
.and_then(|entry| entry["id"].as_str())
|
||||
.map(|s| s.to_string())
|
||||
.ok_or_else(|| format!("No model found in response from {}", models_url))
|
||||
}
|
||||
|
||||
/// Fetch gRPC metadata (returns labels and detected runtime type).
|
||||
async fn fetch_grpc_metadata(
|
||||
url: &str,
|
||||
@@ -283,6 +322,15 @@ impl StepExecutor<LocalWorkerWorkflowData> for DiscoverMetadataStep {
|
||||
}
|
||||
}
|
||||
|
||||
// If no model name discovered yet, try /v1/models as fallback
|
||||
if !labels.contains_key("model_path") && !labels.contains_key("served_model_name") {
|
||||
if let Ok(model_name) =
|
||||
get_model_name_from_v1_models(&config.url, config.api_key.as_deref()).await
|
||||
{
|
||||
labels.insert("served_model_name".to_string(), model_name);
|
||||
}
|
||||
}
|
||||
|
||||
Ok((labels, None))
|
||||
}
|
||||
ConnectionMode::Grpc { .. } => {
|
||||
|
||||
@@ -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(());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user