use std::path::PathBuf;
use std::sync::Arc;
use crate::{
discovery::{ModelManager, ModelWatcher},
engines::StreamingEngineAdapter,
entrypoint::{EngineConfig, RouterConfig, input::common},
grpc::service::kserve,
http::service::metrics::Metrics,
local_model::runtime_config::TokenizerBackend,
namespace::NamespaceFilter,
types::openai::{
chat_completions::{NvCreateChatCompletionRequest, NvCreateChatCompletionStreamResponse},
completions::{NvCreateCompletionRequest, NvCreateCompletionResponse},
},
};
use dynamo_runtime::DistributedRuntime;
pub async fn run(
distributed_runtime: DistributedRuntime,
engine_config: EngineConfig,
) -> anyhow::Result<()> {
let mut grpc_service_builder = kserve::KserveService::builder()
.port(engine_config.local_model().http_port()) .metrics_prefix(engine_config.local_model().metrics_prefix())
.http_cancel_token(Some(distributed_runtime.primary_token()))
.with_request_template(engine_config.local_model().request_template());
if let Some(http_metrics_port) = engine_config.local_model().http_metrics_port() {
grpc_service_builder = grpc_service_builder.http_metrics_port(http_metrics_port);
}
let grpc_service = match engine_config {
EngineConfig::Dynamic {
ref model,
ref prefill_load_estimator,
..
} => {
let grpc_service = grpc_service_builder.build()?;
let router_config = model.router_config();
let migration_limit = model.migration_limit();
let migration_max_seq_len = model.migration_max_seq_len();
let namespace_filter = NamespaceFilter::from_namespace_and_prefix(
model.namespace(),
model.namespace_prefix(),
);
let local_model_path =
(!model.path().as_os_str().is_empty()).then(|| model.path().to_path_buf());
run_watcher(
distributed_runtime.clone(),
grpc_service.state().manager_clone(),
router_config.clone(),
migration_limit,
migration_max_seq_len,
namespace_filter,
prefill_load_estimator.clone(),
local_model_path,
model.metrics_prefix(),
model.runtime_config().tokenizer_backend,
model.runtime_config().tokenizer_fallback_enabled,
)
.await?;
grpc_service
}
EngineConfig::InProcessText { engine, model, .. } => {
let grpc_service = grpc_service_builder.build()?;
let engine = Arc::new(StreamingEngineAdapter::new(engine));
let manager = grpc_service.model_manager();
let checksum = model.card().mdcsum();
manager.add_completions_model(model.service_name(), checksum, engine.clone())?;
manager.add_chat_completions_model(model.service_name(), checksum, engine)?;
grpc_service
}
EngineConfig::InProcessTokens {
engine: inner_engine,
model,
..
} => {
let grpc_service = grpc_service_builder.build()?;
let manager = grpc_service.model_manager();
let checksum = model.card().mdcsum();
let tokenizer = model.card().tokenizer()?;
let chat_pipeline = common::build_pipeline::<
NvCreateChatCompletionRequest,
NvCreateChatCompletionStreamResponse,
>(model.card(), inner_engine.clone(), tokenizer.clone())
.await?;
manager.add_chat_completions_model(model.service_name(), checksum, chat_pipeline)?;
let cmpl_pipeline = common::build_pipeline::<
NvCreateCompletionRequest,
NvCreateCompletionResponse,
>(model.card(), inner_engine, tokenizer)
.await?;
manager.add_completions_model(model.service_name(), checksum, cmpl_pipeline)?;
grpc_service
}
};
let http_service = grpc_service.http_service().clone();
let shutdown_token = distributed_runtime.primary_token();
let join_result = tokio::try_join!(
grpc_service.run(shutdown_token.clone()),
http_service.run(shutdown_token)
);
distributed_runtime.shutdown();
join_result?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
async fn run_watcher(
runtime: DistributedRuntime,
model_manager: Arc<ModelManager>,
router_config: RouterConfig,
migration_limit: u32,
migration_max_seq_len: Option<u32>,
namespace_filter: NamespaceFilter,
prefill_load_estimator: Option<Arc<dyn dynamo_kv_router::PrefillLoadEstimator>>,
local_model_path: Option<PathBuf>,
metrics_prefix: Option<String>,
tokenizer_backend: Option<TokenizerBackend>,
tokenizer_fallback_enabled: Option<bool>,
) -> anyhow::Result<()> {
if crate::lora::lora_serving_enabled() {
let _controller_handle = model_manager.start_lora_controller(runtime.primary_token());
}
let metrics = Arc::new(Metrics::new_with_prefix(metrics_prefix));
let mut watch_obj = ModelWatcher::new(
runtime.clone(),
model_manager,
router_config,
migration_limit,
migration_max_seq_len,
None,
prefill_load_estimator,
metrics,
);
watch_obj.set_local_model_path(local_model_path);
watch_obj.set_tokenizer_backend(tokenizer_backend);
watch_obj.set_tokenizer_fallback_enabled(tokenizer_fallback_enabled);
tracing::debug!("Waiting for remote model");
let discovery = runtime.discovery();
let discovery_stream = discovery
.list_and_watch(
dynamo_runtime::discovery::DiscoveryQuery::AllModels,
Some(runtime.primary_token()),
)
.await?;
let watch_obj = Arc::new(watch_obj);
let _watcher_task = tokio::spawn(async move {
watch_obj.watch(discovery_stream, namespace_filter).await;
});
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{engines::make_echo_engine, local_model::LocalModelBuilder};
use dynamo_runtime::{Runtime, distributed::DistributedConfig};
use std::time::Duration;
#[tokio::test]
async fn metrics_bind_failure_shuts_down_dynamic_and_in_process_runtimes() {
for dynamic in [true, false] {
let occupied = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let model = Box::new(
LocalModelBuilder::default()
.model_name(Some("bind-failure".to_string()))
.http_port(0)
.http_metrics_port(Some(occupied.local_addr().unwrap().port()))
.build()
.await
.unwrap(),
);
let engine_config = if dynamic {
EngineConfig::Dynamic {
model,
chat_engine_factory: None,
prefill_load_estimator: None,
}
} else {
EngineConfig::InProcessText {
engine: make_echo_engine(),
model,
}
};
let drt = DistributedRuntime::new(
Runtime::from_current().unwrap(),
DistributedConfig::process_local(),
)
.await
.unwrap();
let shutdown = drt.primary_token();
let error = tokio::time::timeout(Duration::from_secs(5), run(drt, engine_config))
.await
.expect("gRPC run must return after a metrics-port bind failure")
.expect_err("the occupied metrics port must prevent server startup");
assert!(
error.to_string().contains("already in use"),
"expected a metrics bind error (dynamic={dynamic}), got {error:#}"
);
tokio::time::timeout(Duration::from_secs(5), shutdown.cancelled())
.await
.expect("metrics bind failure must initiate runtime shutdown");
}
}
}