use crate::models::{
ModelInfo, ModelsListResponse, RegisterRequest, ResponseStatus, ServerResponse,
};
use crate::state::{AppState, EndpointModelsCache, ProxyServer, ServerMetrics};
use axum::{extract::State, http::StatusCode, response::IntoResponse, Json};
use hyper_util::client::legacy::Client;
use tracing;
pub async fn register_server(
State(state): State<AppState>,
Json(payload): Json<RegisterRequest>,
) -> impl IntoResponse {
if payload.addr.trim().is_empty() || !payload.addr.contains(':') {
tracing::warn!(
"Invalid address provided for registration: {}",
payload.addr
);
return (
StatusCode::BAD_REQUEST,
Json(ServerResponse {
status: ResponseStatus::Error,
message: "Invalid address format. Expected host:port".to_string(),
}),
);
}
let server_addr = payload.addr.trim().to_string();
tracing::info!("Discovering models from endpoint: {}", server_addr);
let discovered_models = match discover_models(&state.http_client, &server_addr).await {
Ok(models) => models,
Err(e) => {
tracing::error!("Discovery failed for {}: {}", server_addr, e);
return (
StatusCode::BAD_GATEWAY,
Json(ServerResponse {
status: ResponseStatus::Error,
message: format!("Failed to discover models: {}", e),
}),
);
}
};
let mut servers = state.servers.lock().await;
let mut models_cache = state.models_cache.lock().await;
let mut newly_registered = Vec::new();
let mut already_registered = Vec::new();
for model in &discovered_models {
let model_name = model.id.clone();
if servers
.iter()
.any(|s| s.model_name == model_name && s.addr == server_addr)
{
already_registered.push(model_name);
continue;
}
servers.push(ProxyServer {
model_name: model_name.clone(),
addr: server_addr.clone(),
metrics: ServerMetrics::default(),
});
newly_registered.push(model_name.clone());
tracing::info!("Registered model '{}' at {}", model_name, server_addr);
}
if let Some(cache_entry) = models_cache.iter_mut().find(|c| c.endpoint == server_addr) {
cache_entry.models = discovered_models.clone();
cache_entry.discovered_at = std::time::SystemTime::now();
} else {
models_cache.push(EndpointModelsCache {
endpoint: server_addr.clone(),
models: discovered_models.clone(),
discovered_at: std::time::SystemTime::now(),
});
}
drop(servers);
drop(models_cache);
let message = if !newly_registered.is_empty() {
if already_registered.is_empty() {
format!(
"Registered {} model(s): {}",
newly_registered.len(),
newly_registered.join(", ")
)
} else {
format!(
"Registered {} new model(s): {} (Already registered: {})",
newly_registered.len(),
newly_registered.join(", "),
already_registered.join(", ")
)
}
} else {
format!(
"All {} model(s) already registered: {}",
already_registered.len(),
already_registered.join(", ")
)
};
let (status_code, response_status) = if !newly_registered.is_empty() {
(StatusCode::CREATED, ResponseStatus::Success)
} else {
(StatusCode::OK, ResponseStatus::Warning)
};
(
status_code,
Json(ServerResponse {
status: response_status,
message,
}),
)
}
async fn discover_models(
_http_client: &Client<hyper_util::client::legacy::connect::HttpConnector, axum::body::Body>,
addr: &str,
) -> Result<Vec<ModelInfo>, String> {
let url = format!("http://{}/v1/models", addr);
let reqwest_client = reqwest::Client::new();
let response = reqwest_client
.get(&url)
.send()
.await
.map_err(|e| format!("Failed to connect: {}", e))?;
if !response.status().is_success() {
return Err(format!("Endpoint returned HTTP {}", response.status()));
}
let models_response: ModelsListResponse = response
.json()
.await
.map_err(|e| format!("Invalid JSON response: {}", e))?;
if models_response.data.is_empty() {
return Err("No models available at this endpoint".to_string());
}
Ok(models_response.data)
}