mod dispatch;
mod errors;
#[cfg(feature = "pcap")]
mod jobs;
mod limits;
mod operations;
mod protocol;
mod schemas;
mod targets;
mod tools;
use std::sync::{Arc, Mutex};
use tokio::io::{self, AsyncBufReadExt, AsyncWriteExt, BufReader, BufWriter};
use tokio::sync::{mpsc, Semaphore};
use tokio::task::JoinSet;
use dispatch::{handle_request, ServerState};
use protocol::{JsonRpcRequest, JsonRpcResponse};
const MAX_CONCURRENT_REQUESTS: usize = 16;
const MAX_REQUEST_DURATION: std::time::Duration = std::time::Duration::from_secs(15 * 60);
const MAX_REQUEST_LINE_BYTES: usize = 4 * 1024 * 1024;
pub use tools::tools_list;
fn init_tracing() {
use tracing_subscriber::{fmt, EnvFilter};
let filter = EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info"));
let _ = fmt()
.with_writer(std::io::stderr)
.with_env_filter(filter)
.with_target(false)
.json()
.with_current_span(false)
.with_span_list(false)
.try_init();
}
pub async fn run_server() -> anyhow::Result<()> {
init_tracing();
tracing::info!(
version = env!("CARGO_PKG_VERSION"),
pcap = cfg!(feature = "pcap"),
"netscli MCP server starting"
);
let stdin = io::stdin();
let stdout = io::stdout();
let mut reader = BufReader::new(stdin);
let state = Arc::new(Mutex::new(ServerState::default()));
let limiter = Arc::new(Semaphore::new(MAX_CONCURRENT_REQUESTS));
let (tx, mut rx) = mpsc::unbounded_channel::<String>();
let writer_task = tokio::spawn(async move {
let mut writer = BufWriter::new(stdout);
let mut written: u64 = 0;
while let Some(line) = rx.recv().await {
if writer.write_all(line.as_bytes()).await.is_err()
|| writer.write_all(b"\n").await.is_err()
|| writer.flush().await.is_err()
{
break;
}
written += 1;
}
written
});
let send = |tx: &mpsc::UnboundedSender<String>, response: &JsonRpcResponse| {
match serde_json::to_string(response) {
Ok(s) => {
let _ = tx.send(s);
}
Err(e) => tracing::error!(error = %e, "failed to serialize response"),
}
};
let mut handlers: JoinSet<()> = JoinSet::new();
loop {
let mut line = String::new();
match read_bounded_line(&mut reader, &mut line).await {
Ok(0) => break,
Ok(_) => {}
Err(LineError::Invalid(reason)) => {
tracing::warn!(reason, "unreadable line on stdin");
send(&tx, &JsonRpcResponse::parse_error());
continue;
}
Err(LineError::Fatal(e)) => return Err(e.into()),
}
let line = line.trim_end().to_string();
if line.trim().is_empty() {
continue;
}
while handlers.try_join_next().is_some() {}
let request = match serde_json::from_str::<JsonRpcRequest>(&line) {
Ok(req) => req,
Err(e) => {
tracing::warn!(error = %e, "json parse error");
send(&tx, &JsonRpcResponse::parse_error());
continue;
}
};
if request.id.is_none() {
tracing::debug!(method = %request.method, "notification");
if request.method == "notifications/initialized" {
if let Ok(mut guard) = state.lock() {
guard.initialized = true;
}
}
continue;
}
let Ok(permit) = Arc::clone(&limiter).acquire_owned().await else {
break;
};
let state = Arc::clone(&state);
let tx = tx.clone();
let id = request.id.clone().flatten();
handlers.spawn(async move {
let _permit = permit;
let response =
match tokio::time::timeout(MAX_REQUEST_DURATION, handle_request(state, request))
.await
{
Ok(response) => response,
Err(_) => {
tracing::warn!(
timeout_s = MAX_REQUEST_DURATION.as_secs(),
"request exceeded the maximum duration"
);
JsonRpcResponse::request_timeout(id, MAX_REQUEST_DURATION.as_secs())
}
};
match serde_json::to_string(&response) {
Ok(s) => {
let _ = tx.send(s);
}
Err(e) => tracing::error!(error = %e, "failed to serialize response"),
}
});
}
let aborted = handlers.len();
handlers.shutdown().await;
drop(tx);
let requests_handled = writer_task.await.unwrap_or(0);
tracing::info!(
requests_handled,
aborted,
"netscli MCP server shutting down"
);
Ok(())
}
enum LineError {
Invalid(&'static str),
Fatal(std::io::Error),
}
async fn read_bounded_line<R>(reader: &mut R, out: &mut String) -> Result<usize, LineError>
where
R: tokio::io::AsyncBufRead + Unpin,
{
let mut raw = Vec::new();
let mut limited = tokio::io::AsyncReadExt::take(reader, MAX_REQUEST_LINE_BYTES as u64 + 1);
let read = limited
.read_until(b'\n', &mut raw)
.await
.map_err(LineError::Fatal)?;
if read == 0 {
return Ok(0);
}
if raw.len() > MAX_REQUEST_LINE_BYTES {
return Err(LineError::Invalid("line exceeds the maximum request size"));
}
match String::from_utf8(raw) {
Ok(text) => {
out.push_str(&text);
Ok(read)
}
Err(_) => Err(LineError::Invalid("line is not valid UTF-8")),
}
}