mod cli;
use std::sync::Arc;
#[cfg(feature = "profiling")]
#[global_allocator]
static ALLOC: dhat::Alloc = dhat::Alloc;
use clap::{Parser, Subcommand};
use tokio::sync::Mutex;
use sqlserver_mcp_catalog::auth::auth_manager::AuthManager;
use sqlserver_mcp_catalog::core::component_registry::ComponentRegistry;
use sqlserver_mcp_catalog::core::config_manager::load_config;
use sqlserver_mcp_catalog::core::config_schema::Transport;
use sqlserver_mcp_catalog::core::health_check_manager::HealthCheckManager;
use sqlserver_mcp_catalog::core::logger::init_logging;
use sqlserver_mcp_catalog::core::mcp_server::{McpifyServer, connect_stdio};
use sqlserver_mcp_catalog::core::otel;
use sqlserver_mcp_catalog::core::shutdown_handler::{install_shutdown_handlers, on_shutdown};
use sqlserver_mcp_catalog::data::store::{open_store, resolve_store_path};
use sqlserver_mcp_catalog::http::server::start_http_server;
use sqlserver_mcp_catalog::http::types::HttpServerConfig;
#[derive(Parser)]
#[command(
name = "sqlserver-mcp",
about = "SQL Server 2025 - master/msdb/sandbox combined catalog MCP server",
version
)]
struct Cli {
#[command(subcommand)]
command: Command,
}
#[derive(Subcommand)]
enum Command {
Setup,
Search {
query: String,
#[arg(short = 'l', long, default_value_t = 5)]
limit: usize,
#[arg(long, hide = true, default_value_t = 0)]
profile_warmups: usize,
#[arg(long, hide = true, default_value_t = 1)]
profile_iterations: usize,
},
Get { operation_id: String },
Call {
operation_id: String,
#[arg(short = 'a', long, default_value = "{}")]
args: String,
},
Start,
Http {
#[arg(long)]
host: Option<String>,
#[arg(long)]
port: Option<u16>,
#[arg(long)]
cors_allow: Option<String>,
},
TestConnection,
Config,
Version,
Versions,
}
async fn run_harness_server() -> anyhow::Result<()> {
let config = load_config(serde_json::Map::new())?;
let (otel_layer, otel_provider) = match otel::build_layer("sqlserver-mcp") {
Ok((layer, provider)) => (Some(layer), Some(provider)),
Err(_) => (None, None),
};
init_logging(otel_layer);
install_shutdown_handlers();
if let Some(provider) = otel_provider {
on_shutdown(Box::new(move || {
let provider = provider.clone();
Box::pin(async move { otel::shutdown_tracing(provider) })
}));
}
let registry = Arc::new(Mutex::new(ComponentRegistry::new()));
let mut health_checks = HealthCheckManager::new(registry.clone());
let db_path = resolve_store_path(&config.api_version)?;
{
let db_path = db_path.clone();
health_checks
.register(
"store",
true,
Box::new(move || {
let db_path = db_path.clone();
Box::pin(async move {
open_store(&db_path)?;
Ok(())
})
}),
)
.await;
}
let health_checks = Arc::new(health_checks);
let health_check_handle = health_checks.start();
on_shutdown(Box::new(move || {
health_check_handle.abort();
Box::pin(async move {})
}));
let auth_manager = Arc::new(Mutex::new(AuthManager::new(config.auth_method)));
if config.transport == Transport::Http {
let http_config = HttpServerConfig {
host: config.host.clone(),
port: config.port,
cors_allow: config.cors_allow.clone(),
};
let factory_config = config.clone();
start_http_server(
move || {
Ok(McpifyServer::new(
factory_config.api_version.clone(),
factory_config.clone(),
auth_manager.clone(),
))
},
&http_config,
registry,
)
.await?;
} else {
let api_version = config.api_version.clone();
connect_stdio(McpifyServer::new(api_version, config, auth_manager)).await?;
}
Ok(())
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
rustls::crypto::aws_lc_rs::default_provider()
.install_default()
.expect("failed to install rustls crypto provider");
#[cfg(feature = "profiling")]
let _dhat_profiler = dhat::Profiler::new_heap();
let cli = Cli::parse();
let result = match cli.command {
Command::Setup => cli::setup::run().await,
Command::Search {
query,
limit,
profile_warmups,
profile_iterations,
} => cli::search::run(&query, limit, profile_warmups, profile_iterations).await,
Command::Get { operation_id } => cli::get::run(&operation_id).await,
Command::Call { operation_id, args } => cli::call::run(&operation_id, &args).await,
Command::Start => {
unsafe {
std::env::set_var("SQLSERVER_TRANSPORT", "stdio");
}
run_harness_server().await
}
Command::Http {
host,
port,
cors_allow,
} => {
unsafe {
std::env::set_var("SQLSERVER_TRANSPORT", "http");
if let Some(host) = &host {
std::env::set_var("SQLSERVER_HOST", host);
}
if let Some(port) = port {
std::env::set_var("SQLSERVER_PORT", port.to_string());
}
if let Some(cors_allow) = &cors_allow {
std::env::set_var("SQLSERVER_CORS_ALLOW", cors_allow);
}
}
run_harness_server().await
}
Command::TestConnection => cli::test_connection::run().await,
Command::Config => cli::config::run(),
Command::Version => cli::version::run(),
Command::Versions => cli::versions::run(),
};
if let Err(err) = result {
eprintln!("{err}");
std::process::exit(1);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn search_accepts_hidden_profiling_workload_controls() {
let cli = Cli::try_parse_from([
"sqlserver-mcp",
"search",
"test query",
"--profile-warmups",
"2",
"--profile-iterations",
"10",
])
.unwrap();
match cli.command {
Command::Search {
query,
limit,
profile_warmups,
profile_iterations,
} => {
assert_eq!(query, "test query");
assert_eq!(limit, 5);
assert_eq!(profile_warmups, 2);
assert_eq!(profile_iterations, 10);
}
_ => panic!("expected search command"),
}
}
}