mod args;
mod backend;
mod families;
mod installed;
mod playground;
mod proxy;
use std::collections::HashSet;
use std::net::SocketAddr;
use std::path::PathBuf;
use std::sync::Arc;
use anyhow::{bail, Context, Result};
use clap::Parser;
use openkind_api::{grpc, http, AppState, AuthConfig};
use openkind_backends::qwen35::{
Qwen35Backend, Qwen35DecisionEngine, Qwen35EngineConfig, SchedulerConfig,
};
use openkind_engine::{DecisionEngine, EngineRegistry, MockEngine};
use openkind_model_store::{default_models_dir, ModelStore};
use openkind_runtime::{peak_resident_bytes, BackendCapabilities};
use tokio::net::TcpListener;
use tonic::transport::Server;
use tracing::info;
use tracing_subscriber::EnvFilter;
use crate::args::{
parse_grpc_addr, resolve_aliases, Args, ArrowArg, PlaygroundArg, Qwen35BackendArg,
};
fn backend_from_arg(backend: Qwen35BackendArg, cuda_device: usize) -> Qwen35Backend {
backend.to_backend(cuda_device)
}
fn load_qwen(
args: &Args,
bundle_root: PathBuf,
checkpoint_root: PathBuf,
tokenizer_path: PathBuf,
) -> Result<Arc<dyn DecisionEngine>> {
let mut scheduler = SchedulerConfig::for_pinned_profile(
SchedulerConfig::LOWEST_MEASURED_SHARED_SAVINGS_RATIO,
args.qwen35_max_tensor_bytes,
)
.with_backend_capabilities(BackendCapabilities::per_lane())
.with_forced_strategy(args.qwen35_execution.into());
if let Some(max_process_bytes) = args.qwen35_max_process_bytes {
scheduler =
scheduler.with_process_memory(openkind_backends::qwen35::ProcessMemoryEnvelope {
observed_resident_bytes: peak_resident_bytes()
.context("read process peak RSS for native admission")?,
forward_scratch_bytes: args.qwen35_scratch_bytes,
allocator_headroom_bytes: args.qwen35_allocator_headroom_bytes,
max_process_bytes,
});
}
crate::backend::load(
args.qwen35_backend,
&checkpoint_root,
args.device_ordinals(),
|backend| {
Ok(Arc::new(
Qwen35DecisionEngine::load(Qwen35EngineConfig {
bundle_root: bundle_root.clone(),
checkpoint_root: checkpoint_root.clone(),
tokenizer_path: tokenizer_path.clone(),
backend: backend_from_arg(backend, args.cuda_device),
scheduler: scheduler.clone(),
max_concurrent_requests: args.qwen35_concurrency,
max_queued_requests: args.qwen35_queue,
retry_after_ms: 1_000,
evaluation_timeout: Some(std::time::Duration::from_millis(
args.qwen35_timeout_ms,
)),
})
.context("load native Qwen3.5 engine")?,
) as Arc<dyn DecisionEngine>)
},
)
}
fn main() -> Result<()> {
let args = parse_args_from(std::env::args_os())?;
if let Some(runtime) = &args.onnx_runtime {
std::env::set_var("ORT_DYLIB_PATH", runtime);
}
tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()
.context("build tokio runtime")?
.block_on(run(args))
}
fn parse_args_from(
arguments: impl IntoIterator<Item = impl Into<std::ffi::OsString>>,
) -> Result<Args> {
let arguments: Vec<std::ffi::OsString> = arguments.into_iter().map(Into::into).collect();
std::thread::Builder::new()
.name("openkindd-args".into())
.stack_size(2 * 1024 * 1024)
.spawn(move || Args::parse_from(arguments))
.context("start argument parser")?
.join()
.map_err(|_| anyhow::anyhow!("argument parser panicked"))
}
async fn run(args: Args) -> Result<()> {
if let Some(probe) = args.probe_backend {
return backend::probe(probe, args.device_ordinals());
}
if args.diagnose_backends {
return backend::diagnose(args.device_ordinals(), args.json);
}
let api_key = args.resolve_api_key()?;
let mut aliases = HashSet::new();
for alias in &args.models {
if !aliases.insert(alias) {
anyhow::bail!("duplicate --models alias `{alias}`");
}
}
for name in &args.installed_models {
if !aliases.insert(name) {
if args.models.contains(name) {
anyhow::bail!("installed model alias `{name}` collides with --models");
}
anyhow::bail!("duplicate --installed-models alias `{name}`");
}
}
if args.proxy_cache_upstream_key.is_some() {
if args
.proxy_cache_upstream_key
.as_deref()
.is_none_or(|key| key.trim().is_empty())
{
anyhow::bail!("--proxy-cache-upstream-key must be nonempty");
}
if api_key.as_deref().is_none_or(|key| key.trim().is_empty()) {
anyhow::bail!("--proxy-cache-upstream-key requires a nonempty local --api-key");
}
}
if args.proxy_cache_upstream.is_some() {
for alias in &args.proxy_cache_models {
if !aliases.insert(alias) {
anyhow::bail!("proxy-cache alias `{alias}` collides with a local or installed model, or is duplicated");
}
}
}
init_tracing(&args.log_filter)?;
let http_addr = resolve_aliases(&[
("OPENKIND_HTTP_ADDR", args.http_addr),
("OPENDECISION_HTTP_ADDR", args.opendecision_http_addr),
("OPENPICK_HTTP_ADDR", args.legacy_http_addr),
])?
.unwrap_or_else(|| {
"127.0.0.1:8080"
.parse()
.expect("valid default HTTP address")
});
validate_playground_bind(args.playground, http_addr)?;
let grpc_addr_value = resolve_aliases(&[
("OPENKIND_GRPC_ADDR", args.grpc_addr.clone()),
(
"OPENDECISION_GRPC_ADDR",
args.opendecision_grpc_addr.clone(),
),
("OPENPICK_GRPC_ADDR", args.legacy_grpc_addr.clone()),
])?
.unwrap_or_else(|| "127.0.0.1:9090".to_owned());
let grpc_addr = parse_grpc_addr(&grpc_addr_value)
.context("invalid --grpc-addr (expected host:port, or `0` to disable)")?;
let auth = AuthConfig::new(api_key);
if auth.is_required() {
info!("api key auth: enabled (gate on /v1/* and gRPC)");
} else {
if !http_addr.ip().is_loopback()
|| grpc_addr.is_some_and(|grpc_addr| !grpc_addr.ip().is_loopback())
{
tracing::warn!(
"SECURITY WARNING: Server is binding to a non-loopback interface without authentication! Anyone with network access can execute inference queries."
);
} else {
info!("api key auth: disabled (neither OPENKIND_API_KEY nor TYPESAFE_API_KEY set)");
}
}
for accelerator in openkind_runtime::detect_accelerators() {
info!(
device = %accelerator.device,
name = %accelerator.name,
total_memory_bytes = ?accelerator.total_memory_bytes,
"detected accelerator"
);
}
let support = openkind_backends::device::execution_support();
info!(
candle_cuda = support.candle_cuda,
onnx = support.onnx,
onnx_cuda = support.onnx_cuda,
onnx_rocm = support.onnx_rocm,
mlx = support.mlx,
cuda_device = args.cuda_device,
rocm_device = args.rocm_device,
"execution backend support compiled into this daemon"
);
info!(
http = %http_addr,
grpc = ?grpc_addr,
models = ?args.models,
playground = matches!(args.playground, PlaygroundArg::On),
"starting openkindd"
);
let mut registry = EngineRegistry::new();
let mock = Arc::new(MockEngine::new());
args.family_args.validate(&args.qwen35_aliases)?;
let claimed_family_aliases = args.family_args.claimed_aliases()?;
let family_engines = args
.family_args
.load_requested(&args.models, args.device_ordinals())?;
for (alias, engine) in &family_engines {
info!(
alias,
backend = engine.backend_id(),
"registered family engine"
);
}
let native_requested = args
.models
.iter()
.any(|alias| args.qwen35_aliases.contains(alias));
let native: Option<Arc<dyn DecisionEngine>> = if native_requested {
let bundle_root = args
.qwen35_bundle_root
.clone()
.context("native alias requested but --qwen35-bundle-root is missing")?;
let checkpoint_root = args
.qwen35_checkpoint_root
.clone()
.context("native alias requested but --qwen35-checkpoint-root is missing")?;
let tokenizer_path = args
.qwen35_tokenizer
.clone()
.context("native alias requested but --qwen35-tokenizer is missing")?;
Some(load_qwen(
&args,
bundle_root,
checkpoint_root,
tokenizer_path,
)?)
} else {
None
};
for alias in &args.models {
if args.family_args.winnow_aliases.contains(alias)
|| args.family_args.router_script_aliases.contains(alias)
{
continue;
}
let engine: Arc<dyn DecisionEngine> = if args.qwen35_aliases.contains(alias) {
native
.as_ref()
.expect("native engine loaded when a native alias is requested")
.clone()
} else if let Some((_, family_engine)) = family_engines
.iter()
.find(|(family_alias, _)| family_alias == alias)
{
family_engine.clone()
} else if claimed_family_aliases.contains(alias) {
bail!("alias `{alias}` is configured for a family engine but was not loaded");
} else {
mock.clone()
};
info!(alias, backend = engine.backend_id(), "registered model");
registry.register(alias.clone(), engine);
}
for (alias, sibling_labels) in args.family_args.winnow_requested(&args.models)? {
let (model_root, adapter) = args.family_args.winnow_artifacts(&alias)?;
let mut siblings: Vec<(String, Arc<dyn DecisionEngine>)> = Vec::new();
for sibling_alias in &sibling_labels {
let engine = registry.get(sibling_alias).ok_or_else(|| {
anyhow::anyhow!(
"winnow alias `{alias}` references unregistered sibling `{sibling_alias}`"
)
})?;
siblings.push((sibling_alias.clone(), engine));
}
let engine = backend::load(
args.family_args.winnow_backend,
&model_root,
args.device_ordinals(),
|backend| {
openkind_backends::families::winnow::WinnowEngine::load_with_execution(
openkind_backends::families::winnow::WinnowEngineConfig {
model_root: model_root.clone(),
adapter_path: adapter.clone(),
limits: openkind_backends::families::support::FamilyLimits {
max_concurrent_requests: args.family_args.family_concurrency,
max_queued_requests: args.family_args.family_queue,
retry_after_ms: 1_000,
evaluation_timeout: Some(std::time::Duration::from_millis(
args.family_args.family_timeout_ms,
)),
},
},
siblings.clone(),
backend.to_execution(args.device_ordinals())?,
)
.with_context(|| format!("compose winnow alias `{alias}`"))
},
)?;
info!(alias, backend = engine.backend_id(), "registered model");
registry.register(alias, Arc::new(engine));
}
for (alias, rules) in args.family_args.router_script_requested(&args.models)? {
let siblings: std::collections::HashMap<String, Arc<dyn DecisionEngine>> = rules
.referenced_aliases()
.into_iter()
.filter_map(|sibling_alias| {
registry
.get(sibling_alias)
.map(|engine| (sibling_alias.to_owned(), engine))
})
.collect();
let router = openkind_backends::families::router_script::RouterScriptEngine::new(
rules, &siblings,
)
.map_err(|error| anyhow::anyhow!("compose router-script alias `{alias}`: {error}"))?;
info!(alias, backend = router.backend_id(), "registered model");
registry.register(alias, Arc::new(router));
}
let mut _installed_guards = Vec::new();
if !args.installed_models.is_empty() {
let dir = match &args.models_dir {
Some(dir) => dir.clone(),
None => default_models_dir()?,
};
let store = ModelStore::new(dir)?;
let mut deferred_winnow = Vec::new();
for name in &args.installed_models {
let installed = store.acquire_serving(name)?;
let kind = installed::installed_kind(&installed.manifest)
.with_context(|| format!("unsupported installed model profile `{name}`"))?;
if matches!(kind, installed::InstalledKind::Winnow) {
deferred_winnow.push((name.clone(), installed));
continue;
}
let engine = installed::load_installed_engine(&args, kind, &installed.root, ®istry)?;
info!(
alias = name,
backend = engine.backend_id(),
"registered installed model"
);
registry.register(name.clone(), engine);
_installed_guards.push(installed);
}
for (name, installed) in deferred_winnow {
let engine = installed::load_installed_engine(
&args,
installed::InstalledKind::Winnow,
&installed.root,
®istry,
)?;
info!(
alias = name,
backend = engine.backend_id(),
"registered installed model"
);
registry.register(name, engine);
_installed_guards.push(installed);
}
}
let mut _proxy_encoder_guard: Option<openkind_model_store::InstalledModel> = None;
let proxy_service: Option<Arc<proxy::ProxyService>> = match &args.proxy_cache_upstream {
Some(upstream) => {
if !(0.0..1.0).contains(&args.proxy_cache_target_agreement) {
anyhow::bail!(
"--proxy-cache-target-agreement must be in (0, 1), got {}",
args.proxy_cache_target_agreement
);
}
let (embedder, guard) = proxy::resolve_encoder(
&args.proxy_cache_encoder,
args.proxy_cache_encoder_backend,
args.cuda_device,
args.models_dir.as_deref(),
)
.await?;
_proxy_encoder_guard = guard;
let task_config = openkind_backends::proxy_cache::TaskConfig {
store_text: args.proxy_cache_store_text,
min_train_samples: args.proxy_cache_min_train_samples,
min_calib_samples: args.proxy_cache_min_calib_samples,
shadow_min_samples: args.proxy_cache_shadow_min_samples,
calib_fraction: args.proxy_cache_calib_fraction,
min_new_samples: args.proxy_cache_min_new_samples,
..Default::default()
};
let data_dir = match &args.proxy_cache_data_dir {
Some(dir) => dir.clone(),
None => proxy::default_proxy_cache_dir()
.map_err(|error| anyhow::anyhow!("proxy cache data dir: {error}"))?,
};
let manager = openkind_backends::proxy_cache::ProxyCacheManager::new(
openkind_backends::proxy_cache::ProxyCacheManagerConfig {
data_dir,
task_config,
target_agreement: args.proxy_cache_target_agreement,
confidence_floor: None,
admission_min_requests: args.proxy_cache_admission_min,
..Default::default()
},
embedder,
)
.map_err(|error| anyhow::anyhow!("proxy cache manager: {error}"))?;
let service = Arc::new(proxy::ProxyService::new(
manager,
proxy::ProxyCacheServiceConfig {
upstream: upstream.clone(),
upstream_key: args.proxy_cache_upstream_key.clone(),
upstream_timeout_ms: args.proxy_cache_upstream_timeout_ms,
proxied_models: args.proxy_cache_models.clone(),
},
));
for alias in &args.proxy_cache_models {
registry.register(
alias.clone(),
Arc::new(proxy::ProxyForwardEngine::new(
service.clone(),
alias.clone(),
)),
);
info!(
alias,
backend = "proxy-cache/upstream-forward",
"registered proxied model (gRPC forwards upstream; HTTP answers via the cache hook)"
);
}
info!(
upstream = %upstream,
encoder = %args.proxy_cache_encoder,
models = ?args.proxy_cache_models,
"proxy cache enabled"
);
Some(service)
}
None => None,
};
openkind_api::http::install_metrics_recorder().context("metrics recorder")?;
let args = Arc::new(args);
let mut state = AppState::new(registry);
state.proxy = proxy_service.map(|service| service as Arc<dyn openkind_api::SystemProxy>);
if matches!(args.playground, PlaygroundArg::On) {
state.playground_models = Some(Arc::new(playground::LocalModels::new(
args.clone(),
state.registry.clone(),
_installed_guards,
)?));
}
let (shutdown_tx, mut shutdown_rx_http) = tokio::sync::watch::channel(false);
let mut shutdown_rx_grpc = shutdown_tx.subscribe();
let sig_tx = shutdown_tx.clone();
tokio::spawn(shutdown_signal(sig_tx));
let http_state = state.clone();
let http_auth = auth.clone();
let http_tx = shutdown_tx.clone();
let limits = openkind_api::RequestLimits::new(openkind_api::RateLimitConfig {
max_requests: args.rate_limit_rpm,
window: std::time::Duration::from_secs(60),
});
let http_limits = limits.clone();
let playground_enabled = matches!(args.playground, PlaygroundArg::On);
let arrow_enabled = matches!(args.arrow, ArrowArg::On);
let http_handle = tokio::spawn(async move {
let _shutdown = ShutdownOnDrop(http_tx);
let router = http::router_daemon_with_arrow_and_limits(
http_state,
http_auth,
openkind_api::http::MAX_PAYLOAD_SIZE_BYTES,
http_limits,
playground_enabled,
arrow_enabled,
);
let listener = TcpListener::bind(http_addr)
.await
.with_context(|| format!("bind http {http_addr}"))?;
info!(http_addr = %listener.local_addr()?, "http listening");
axum::serve(
listener,
router.into_make_service_with_connect_info::<SocketAddr>(),
)
.with_graceful_shutdown(async move {
let _ = shutdown_rx_http.wait_for(|&v| v).await;
})
.await
.context("http serve")
});
let grpc_state = state.clone();
let grpc_tx = shutdown_tx.clone();
let grpc_handle = if let Some(grpc_addr) = grpc_addr {
let svc = grpc::service_with_auth_and_limits(
(*grpc_state.registry).clone(),
auth.clone(),
limits,
);
Some(tokio::spawn(async move {
let _shutdown = ShutdownOnDrop(grpc_tx);
info!(%grpc_addr, "grpc listening");
Server::builder()
.add_service(svc)
.serve_with_shutdown(grpc_addr, async move {
let _ = shutdown_rx_grpc.wait_for(|&v| v).await;
})
.await
.context("grpc serve")
}))
} else {
info!("grpc disabled (--grpc-addr 0)");
None
};
if let Some(handle) = grpc_handle {
let (h, g) = tokio::join!(http_handle, handle);
h.context("http task")??;
g.context("grpc task")??;
} else {
http_handle
.await
.context("http task")?
.context("http serve")?;
}
info!("openkindd exited cleanly");
Ok(())
}
struct ShutdownOnDrop(tokio::sync::watch::Sender<bool>);
impl Drop for ShutdownOnDrop {
fn drop(&mut self) {
let _ = self.0.send(true);
}
}
fn init_tracing(filter: &str) -> Result<()> {
let env_filter = EnvFilter::try_new(filter).context("invalid log filter")?;
tracing_subscriber::fmt()
.with_env_filter(env_filter)
.with_target(true)
.init();
Ok(())
}
fn validate_playground_bind(playground: PlaygroundArg, http_addr: SocketAddr) -> Result<()> {
if matches!(playground, PlaygroundArg::On) && !http_addr.ip().is_loopback() {
anyhow::bail!(
"--playground on requires a loopback --http-addr (for example 127.0.0.1:8080) because the playground can change loaded models"
);
}
Ok(())
}
async fn shutdown_signal(shutdown: tokio::sync::watch::Sender<bool>) {
#[cfg(unix)]
let mut interrupt = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::interrupt())
.expect("install SIGINT handler");
#[cfg(unix)]
let mut terminate = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
.expect("install SIGTERM handler");
for attempt in 0..2 {
#[cfg(unix)]
let code = tokio::select! {
Some(_) = interrupt.recv() => 130,
Some(_) = terminate.recv() => 143,
};
#[cfg(not(unix))]
let code = {
tokio::signal::ctrl_c()
.await
.expect("install Ctrl-C handler");
130
};
if attempt == 0 {
info!("received shutdown signal; draining listeners");
let _ = shutdown.send(true);
} else {
tracing::warn!("received second shutdown signal; forcing termination");
std::process::exit(code);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn startup_validation_runs_on_a_small_stack() {
let worker = std::thread::Builder::new()
.stack_size(512 * 1024)
.spawn(|| {
let args = parse_args_from([
"openkindd",
"--models",
"mock",
"--installed-models",
"mock",
])
.unwrap();
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
runtime.block_on(run(args)).unwrap_err().to_string()
})
.unwrap();
assert!(worker.join().unwrap().contains("collides with --models"));
}
#[tokio::test]
async fn listener_panic_notifies_peer_shutdown() {
let (tx, mut rx) = tokio::sync::watch::channel(false);
let listener = tokio::spawn(async move {
let _shutdown = ShutdownOnDrop(tx);
panic!("simulated listener failure");
});
assert!(listener.await.unwrap_err().is_panic());
tokio::time::timeout(std::time::Duration::from_secs(1), rx.wait_for(|&v| v))
.await
.unwrap()
.unwrap();
}
#[test]
fn playground_requires_a_loopback_http_listener() {
for address in ["127.0.0.1:8080", "[::1]:8080"] {
validate_playground_bind(PlaygroundArg::On, address.parse().unwrap()).unwrap();
}
let error = validate_playground_bind(PlaygroundArg::On, "0.0.0.0:8080".parse().unwrap())
.unwrap_err();
assert!(error.to_string().contains("requires a loopback"));
validate_playground_bind(PlaygroundArg::Off, "0.0.0.0:8080".parse().unwrap()).unwrap();
}
#[tokio::test]
async fn proxy_validation_precedes_encoder_resolution_and_listeners() {
for arguments in [
vec![
"openkindd",
"--proxy-cache-upstream",
"http://127.0.0.1:1",
"--proxy-cache-upstream-key",
"sponsored",
"--models",
"mock",
],
vec![
"openkindd",
"--proxy-cache-upstream",
"http://127.0.0.1:1",
"--models",
"jev-latest",
],
vec![
"openkindd",
"--proxy-cache-upstream",
"http://127.0.0.1:1",
"--models",
"mock",
"--installed-models",
"jev-latest",
],
] {
let args = parse_args_from(arguments).unwrap();
let error = run(args).await.unwrap_err().to_string();
assert!(
error.contains("requires a nonempty") || error.contains("collides"),
"{error}"
);
}
}
}