use clap::{Parser, ValueEnum};
use meerkat_core::{Config, RealmConfig, RealmSelection, RuntimeBootstrap};
use meerkat_rest::{AppState, router};
use meerkat_store::RealmBackend;
use std::{net::SocketAddr, path::PathBuf};
use tower_http::{cors::CorsLayer, trace::TraceLayer};
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};
#[derive(Parser, Debug)]
#[command(name = "rkat-rest", version = env!("CARGO_PKG_VERSION"))]
struct Args {
#[arg(long)]
realm: Option<String>,
#[arg(long)]
isolated: bool,
#[arg(long)]
instance: Option<String>,
#[arg(long, value_enum)]
realm_backend: Option<RealmBackendArg>,
#[arg(long)]
state_root: Option<PathBuf>,
#[arg(long)]
context_root: Option<PathBuf>,
#[arg(long)]
user_config_root: Option<PathBuf>,
#[arg(long, default_value_t = false)]
expose_paths: bool,
}
#[derive(Clone, Copy, Debug, ValueEnum)]
enum RealmBackendArg {
Jsonl,
Sqlite,
}
impl From<RealmBackendArg> for RealmBackend {
fn from(value: RealmBackendArg) -> Self {
match value {
RealmBackendArg::Jsonl => RealmBackend::Jsonl,
RealmBackendArg::Sqlite => RealmBackend::Sqlite,
}
}
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let args = Args::parse();
let selection = RealmConfig::selection_from_inputs(
args.realm.clone(),
args.isolated,
RealmSelection::Isolated,
)?;
tracing_subscriber::registry()
.with(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| "meerkat_rest=info,tower_http=debug".into()),
)
.with(tracing_subscriber::fmt::layer())
.init();
let state = AppState::load_with_bootstrap_and_options(
RuntimeBootstrap {
realm: RealmConfig {
selection,
instance_id: args.instance,
backend_hint: args
.realm_backend
.map(Into::into)
.map(|b: RealmBackend| b.as_str().to_string()),
state_root: args.state_root,
},
context: meerkat_core::ContextConfig {
context_root: args.context_root,
user_config_root: args.user_config_root,
},
},
args.expose_paths,
)
.await?;
let mut config = state
.config_store
.get()
.await
.unwrap_or_else(|_| Config::default());
if let Err(err) = config.apply_env_overrides() {
tracing::warn!("Failed to apply env overrides: {}", err);
}
tracing::info!(
realm_id = %state.realm,
backend = %state.backend,
store_path = %state.store_path.display(),
default_model = %meerkat::resolve_create_session_default_model(&config),
max_tokens = state.max_tokens,
"Starting Meerkat REST server"
);
let env_var_names: &[&str] = &[
"ANTHROPIC_API_KEY",
"OPENAI_API_KEY",
"GEMINI_API_KEY",
"GOOGLE_API_KEY",
"RKAT_ANTHROPIC_API_KEY",
"RKAT_OPENAI_API_KEY",
"RKAT_GEMINI_API_KEY",
];
let has_api_key =
!config.realm.is_empty() || env_var_names.iter().any(|v| std::env::var(v).is_ok());
if !has_api_key {
tracing::warn!(
"No provider API key configured (config, realm, or environment). \
API calls will fail until a key is set."
);
}
let addr: SocketAddr = format!("{}:{}", state.rest_host, state.rest_port)
.parse()
.map_err(|e| format!("Invalid host:port combination: {e}"))?;
let shutdown_state = state.clone();
let app = router(state)
.layer(TraceLayer::new_for_http())
.layer(CorsLayer::permissive());
tracing::info!("Listening on http://{}", addr);
let listener = tokio::net::TcpListener::bind(addr).await?;
axum::serve(listener, app)
.with_graceful_shutdown(shutdown_signal())
.await?;
shutdown_state
.request_executor
.shutdown_and_abort_stragglers()
.await;
#[cfg(feature = "mcp")]
meerkat_rest::shutdown_all_mcp_sessions(&shutdown_state).await;
shutdown_state.shutdown_schedule_host().await;
tracing::info!("Server shutdown complete");
Ok(())
}
async fn shutdown_signal() {
let ctrl_c = async {
if let Err(e) = tokio::signal::ctrl_c().await {
tracing::error!("Failed to install Ctrl+C handler: {}", e);
}
};
#[cfg(unix)]
let terminate = async {
match tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) {
Ok(mut sig) => {
sig.recv().await;
}
Err(e) => {
tracing::error!("Failed to install signal handler: {}", e);
}
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
() = ctrl_c => {
tracing::info!("Received Ctrl+C, shutting down...");
},
() = terminate => {
tracing::info!("Received SIGTERM, shutting down...");
},
}
}