use super::handlers::{ToolHandler, all_tool_handlers};
use super::protocol::{JsonRpcError, JsonRpcRequest, JsonRpcResponse};
use super::request_meta::{collect_request_timings, elapsed_ms};
use crate::cli::registry::ProjectRegistry;
use anyhow::Context;
use axum::{
Router,
extract::Json,
http::{HeaderMap, HeaderValue, StatusCode},
response::{
IntoResponse, Response,
sse::{Event, KeepAlive, Sse},
},
routing::{get, post},
};
use dashmap::DashMap;
use futures_util::stream::{Stream, StreamExt};
use serde_json::Value;
use std::convert::Infallible;
use std::net::SocketAddr;
use std::sync::{
Arc,
atomic::{AtomicBool, AtomicU64, Ordering},
};
use std::time::Instant;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use tracing::{debug, error, info, warn};
#[cfg(unix)]
const SOCKET_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
#[cfg(unix)]
const INITIAL_SOCKET_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(120);
#[cfg(unix)]
fn socket_read_timeout(first_frame: bool) -> std::time::Duration {
if first_frame {
INITIAL_SOCKET_READ_TIMEOUT
} else {
SOCKET_READ_TIMEOUT
}
}
pub static SERVER_STATE: std::sync::OnceLock<Arc<ProjectRegistry>> = std::sync::OnceLock::new();
pub static SERVER_INSTANCE: std::sync::OnceLock<Arc<McpServer>> = std::sync::OnceLock::new();
pub static HANDLERS: std::sync::OnceLock<Vec<ToolHandler>> = std::sync::OnceLock::new();
static SESSION_COUNTER: AtomicU64 = AtomicU64::new(1);
fn generate_session_id() -> String {
let pid = std::process::id();
let seq = SESSION_COUNTER.fetch_add(1, Ordering::Relaxed);
format!("leindex-{pid}-{seq}")
}
pub const DEFAULT_MCP_PORT: u16 = 47500;
pub const BIND_FALLBACK_PORT_RANGE: u16 = 10;
pub const DEFAULT_MAX_HTTP_SESSIONS: usize = 1000;
pub const MAX_SESSIONS_ENV: &str = "LEINDEX_MAX_SESSIONS";
pub fn max_http_sessions() -> usize {
match std::env::var(MAX_SESSIONS_ENV) {
Ok(v) => v
.trim()
.parse::<usize>()
.ok()
.filter(|n| *n > 0)
.unwrap_or(DEFAULT_MAX_HTTP_SESSIONS),
Err(_) => DEFAULT_MAX_HTTP_SESSIONS,
}
}
#[derive(Clone, Debug)]
pub struct McpServerConfig {
pub bind_address: SocketAddr,
pub enable_cors: bool,
pub max_request_size_mb: usize,
}
impl Default for McpServerConfig {
fn default() -> Self {
Self {
bind_address: SocketAddr::from(([127, 0, 0, 1], DEFAULT_MCP_PORT)),
enable_cors: true,
max_request_size_mb: 10,
}
}
}
#[derive(Clone)]
pub struct McpServer {
pub config: McpServerConfig,
pub _registry: Arc<ProjectRegistry>,
pub(crate) handshake_complete: Arc<AtomicBool>,
pub(crate) session_handshakes: Arc<DashMap<Arc<str>, (bool, Instant)>>,
pub(crate) in_flight: Arc<DashMap<Arc<str>, ()>>,
pub(crate) freshness_advisories: Arc<DashMap<(Arc<str>, std::path::PathBuf), u64>>,
}
impl McpServer {
pub fn new(config: McpServerConfig) -> anyhow::Result<Self> {
let registry = Arc::new(ProjectRegistry::new(
crate::cli::registry::DEFAULT_MAX_PROJECTS,
));
SERVER_STATE
.set(registry.clone())
.map_err(|_| anyhow::anyhow!("Server state already initialized"))?;
let handlers: Vec<ToolHandler> = all_tool_handlers();
HANDLERS
.set(handlers)
.map_err(|_| anyhow::anyhow!("Handlers already initialized"))?;
info!(
"MCP server initialized (multi-project registry, max {} projects)",
crate::cli::registry::DEFAULT_MAX_PROJECTS
);
let server = Self {
config,
_registry: registry,
handshake_complete: Arc::new(AtomicBool::new(false)),
session_handshakes: Arc::new(DashMap::new()),
in_flight: Arc::new(DashMap::new()),
freshness_advisories: Arc::new(DashMap::new()),
};
SERVER_INSTANCE
.set(Arc::new(server.clone()))
.map_err(|_| anyhow::anyhow!("Server instance already initialized"))?;
Ok(server)
}
pub fn with_address(bind_address: SocketAddr) -> anyhow::Result<Self> {
let config = McpServerConfig {
bind_address,
..Default::default()
};
Self::new(config)
}
pub fn cleanup_stale_sessions(&self, max_idle: std::time::Duration) -> usize {
let before = self.session_handshakes.len();
self.session_handshakes.retain(|sid, (_, last_access)| {
if self.in_flight.contains_key(sid.as_ref()) {
true
} else {
last_access.elapsed() < max_idle
}
});
let removed = before - self.session_handshakes.len();
self.freshness_advisories
.retain(|(session_id, _), _| self.session_handshakes.contains_key(session_id.as_ref()));
removed
}
fn apply_freshness_advisory(
&self,
session_id: &str,
project_path: Option<&str>,
result: &mut Value,
) {
let result_project_path = result
.get("project_path")
.and_then(Value::as_str)
.map(str::to_owned);
let Some(freshness) = result
.get_mut("_meta")
.and_then(Value::as_object_mut)
.and_then(|meta| meta.get_mut("freshness"))
.and_then(Value::as_object_mut)
else {
return;
};
let generation = freshness
.get("generation")
.and_then(Value::as_u64)
.unwrap_or(0);
let project = project_path
.filter(|path| !path.is_empty())
.map(std::path::PathBuf::from)
.or_else(|| result_project_path.map(std::path::PathBuf::from))
.unwrap_or_else(|| std::path::PathBuf::from("<current>"));
let key = (Arc::<str>::from(session_id), project);
let already_shown = self
.freshness_advisories
.get(&key)
.is_some_and(|seen| *seen == generation);
let advisory = freshness.remove("warning");
if already_shown {
freshness.insert("advisory".to_string(), Value::Null);
} else {
freshness.insert("advisory".to_string(), advisory.unwrap_or(Value::Null));
self.freshness_advisories.insert(key, generation);
}
}
pub fn active_session_count(&self) -> usize {
self.session_handshakes.len()
}
pub fn session_in_flight(&self, session_id: &str) -> bool {
self.in_flight.contains_key(session_id)
}
pub fn begin_request(&self, session_id: &str) {
let key = self
.session_handshakes
.get(session_id)
.map(|entry| entry.key().clone())
.unwrap_or_else(|| Arc::<str>::from(session_id));
self.in_flight.insert(key, ());
}
pub fn end_request(&self, session_id: &str) {
self.in_flight.remove(session_id);
}
pub fn in_flight_guard(server: &Arc<Self>, session_id: &str) -> InFlightGuard {
server.begin_request(session_id);
InFlightGuard {
server: Arc::clone(server),
session_id: Arc::<str>::from(session_id),
}
}
pub async fn run(self) -> anyhow::Result<()> {
let bind_address = self.config.bind_address;
let listener = bind_with_fallback(bind_address).await?;
if listener.local_addr()? != bind_address {
warn!(
"Default port {} was unavailable; bound to fallback {}",
bind_address.port(),
listener.local_addr()?.port()
);
}
self.serve(listener).await
}
pub async fn serve(self, listener: tokio::net::TcpListener) -> anyhow::Result<()> {
let bind_address = listener.local_addr().unwrap_or(self.config.bind_address);
let router = Self::router();
let cleanup_server = self.clone();
let _cleanup_handle = tokio::spawn(async move {
const CLEANUP_INTERVAL: std::time::Duration = std::time::Duration::from_secs(60);
const SESSION_MAX_IDLE: std::time::Duration = std::time::Duration::from_secs(300); let mut interval = tokio::time::interval(CLEANUP_INTERVAL);
loop {
interval.tick().await;
let removed = cleanup_server.cleanup_stale_sessions(SESSION_MAX_IDLE);
if removed > 0 {
debug!("Cleaned {} stale session(s)", removed);
}
}
});
tokio::spawn(async move {
match _cleanup_handle.await {
Ok(_) => {}
Err(e) => error!("cleanup task died: {e}"),
}
});
info!("Starting MCP server on {}", bind_address);
axum::serve(listener, router.into_make_service())
.await
.context("Server error")?;
Ok(())
}
fn router() -> Router {
Router::new()
.route("/mcp", post(json_rpc_handler))
.route("/mcp/tools/list", get(list_tools_handler))
.route("/health", get(health_check_handler))
.route("/mcp/index/stream", post(index_stream_handler))
}
}
pub(crate) async fn bind_with_fallback(
preferred: SocketAddr,
) -> anyhow::Result<tokio::net::TcpListener> {
if preferred.port() == 0 {
return tokio::net::TcpListener::bind(preferred)
.await
.map_err(|e| anyhow::anyhow!("failed to bind to ephemeral port {}: {}", preferred, e));
}
let mut last_err: Option<std::io::Error> = None;
for offset in 0..=BIND_FALLBACK_PORT_RANGE {
let port = match preferred.port().checked_add(offset) {
Some(p) => p,
None => break,
};
let candidate = SocketAddr::new(preferred.ip(), port);
match tokio::net::TcpListener::bind(candidate).await {
Ok(listener) => return Ok(listener),
Err(e) => {
debug!("bind({}) failed: {}", candidate, e);
last_err = Some(e);
}
}
}
match tokio::net::TcpListener::bind(SocketAddr::new(preferred.ip(), 0)).await {
Ok(listener) => Ok(listener),
Err(ephemeral_err) => {
let preferred_err = last_err
.as_ref()
.map(|e| format!(" ({})", e))
.unwrap_or_default();
Err(anyhow::anyhow!(
"failed to bind to {} or any of the next {} ports{}; \
ephemeral-bind fallback also failed: {}",
preferred,
BIND_FALLBACK_PORT_RANGE,
preferred_err,
ephemeral_err,
))
}
}
}
pub async fn index_stream_handler(
Json(body): Json<Value>,
) -> Sse<impl Stream<Item = Result<Event, Infallible>> + Send> {
use super::protocol::ProgressEvent;
let (tx, rx) = mpsc::channel::<ProgressEvent>(100);
tokio::spawn(async move {
let state = match SERVER_STATE.get() {
Some(s) => s,
None => {
let _ = tx
.send(ProgressEvent::error("Server not initialized"))
.await;
return;
}
};
let project_path = match body.get("project_path").and_then(|v: &Value| v.as_str()) {
Some(p) => p.to_string(),
None => {
let _ = tx.send(ProgressEvent::error("Missing project_path")).await;
return;
}
};
let force_reindex = match body.get("force_reindex") {
Some(Value::Bool(v)) => *v,
Some(Value::String(v)) => {
matches!(v.to_ascii_lowercase().as_str(), "true" | "1" | "yes")
}
Some(Value::Number(v)) => v.as_u64().map(|n| n != 0).unwrap_or(false),
_ => false,
};
let _ = tx
.send(ProgressEvent::progress(
"starting",
0,
0,
format!("Starting indexing for: {}", project_path),
))
.await;
match index_with_progress(state, &project_path, force_reindex, tx.clone()).await {
Ok(stats) => {
let _ = tx
.send(ProgressEvent::complete(
"indexing",
format!("Done: {} files", stats.files_parsed),
))
.await;
}
Err(e) => {
let _ = tx.send(ProgressEvent::error(format!("Error: {}", e))).await;
}
}
});
let stream = ReceiverStream::new(rx).map(|event| -> Result<Event, Infallible> {
let event_data = Event::default()
.json_data(event)
.unwrap_or_else(|_| Event::default().data("error"));
Ok(event_data)
});
Sse::new(stream).keep_alive(
KeepAlive::new()
.interval(std::time::Duration::from_secs(15))
.text("keep-alive"),
)
}
pub async fn index_with_progress(
registry: &Arc<ProjectRegistry>,
project_path: &str,
force_reindex: bool,
tx: mpsc::Sender<super::protocol::ProgressEvent>,
) -> Result<crate::cli::leindex::IndexStats, JsonRpcError> {
use super::protocol::ProgressEvent;
let handle = registry.get_or_load(Some(project_path)).await?;
let cached_stats = {
let idx = handle.read().await;
if idx.is_indexed() && !idx.is_stale_fast() && !force_reindex {
Some(idx.get_stats().clone())
} else {
None
}
};
if let Some(stats) = cached_stats {
let _ = tx
.send(ProgressEvent::progress("skipping", 1, 1, "Already indexed"))
.await;
return Ok(stats);
}
let _ = tx
.send(ProgressEvent::progress(
"collecting",
0,
0,
"Collecting source files...",
))
.await;
let _ = tx
.send(ProgressEvent::progress(
"consolidating",
0,
0,
"Waiting for any in-flight index on this project...",
))
.await;
let stats = registry
.index_project(Some(project_path), force_reindex)
.await?;
let _ = tx
.send(ProgressEvent::progress(
"loading_storage",
0,
0,
"Loading indexed data...",
))
.await;
Ok(stats)
}
fn handle_initialize(server: &McpServer) -> (Value, Option<String>) {
let session_id = generate_session_id();
{
let max_sessions = max_http_sessions();
if server.session_handshakes.len() >= max_sessions {
let oldest_id = server
.session_handshakes
.iter()
.filter(|r| !server.in_flight.contains_key(r.key().as_ref()))
.min_by_key(|r| r.value().1)
.map(|r| r.key().clone());
if let Some(id) = oldest_id {
server.session_handshakes.remove(id.as_ref());
server
.freshness_advisories
.retain(|(session_id, _), _| session_id.as_ref() != id.as_ref());
}
}
server.session_handshakes.insert(
Arc::<str>::from(session_id.as_str()),
(true, Instant::now()),
);
}
let result = serde_json::json!({
"protocolVersion": "2024-11-05",
"capabilities": {
"tools": {
"listChanged": true
},
"prompts": {
"listChanged": true
},
"resources": {
"listChanged": true,
"subscribe": false
},
"logging": {},
"progress": true
},
"serverInfo": {
"name": "leindex",
"version": env!("CARGO_PKG_VERSION"),
"description": "LeIndex MCP Server - Semantic code indexing and analysis with PDG-based tools for superior code comprehension"
},
"instructions": [
"Projects are no longer auto-indexed on startup. Use explicit tool calls to index projects.",
"The server must receive an 'initialize' call before processing other requests."
]
});
(result, Some(session_id))
}
fn handle_ping() -> Value {
serde_json::json!({})
}
async fn json_rpc_handler(headers: HeaderMap, Json(body): Json<Value>) -> Response {
let transport_started = Instant::now();
let incoming_session_id = headers
.get("Mcp-Session-Id")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
let json_req: JsonRpcRequest = match serde_json::from_value(body.clone()) {
Ok(r) => r,
Err(e) => {
warn!("Failed to parse JSON-RPC request: {}", e);
return Json(serde_json::json!({
"jsonrpc": "2.0",
"id": null,
"error": {
"code": -32700,
"message": "Invalid JSON"
}
}))
.into_response();
}
};
let server_instance = match SERVER_INSTANCE.get() {
Some(s) => s,
None => {
warn!("Server instance not initialized");
return Json(serde_json::json!({
"jsonrpc": "2.0",
"id": json_req.id,
"error": {
"code": -32603,
"message": "Server instance not initialized"
}
}))
.into_response();
}
};
let state = server_instance._registry.clone();
let handlers = match HANDLERS.get() {
Some(h) => h,
None => {
warn!("Handlers not initialized");
return Json(serde_json::json!({
"jsonrpc": "2.0",
"id": json_req.id,
"error": {
"code": -32603,
"message": "Handlers not initialized"
}
}))
.into_response();
}
};
debug!("Received JSON-RPC request: method={}", json_req.method);
let id = json_req.id.clone().unwrap_or(serde_json::Value::Null);
if let Err(e) = json_req.validate() {
warn!("Invalid JSON-RPC request: {}", e);
let resp = JsonRpcResponse::error(id, e);
return Json(serde_json::to_value(&resp).unwrap()).into_response();
}
let is_notification = json_req.id.is_none();
if is_notification {
return StatusCode::NO_CONTENT.into_response();
}
if json_req.method == "initialize" {
} else if json_req.method == "ping" {
} else {
let session_ok = match &incoming_session_id {
Some(sid) => {
match server_instance.session_handshakes.get_mut(sid.as_str()) {
Some(mut entry) => {
entry.1 = Instant::now();
entry.0
}
_ => false,
}
}
None => false,
};
if !session_ok {
return Json(serde_json::json!({
"jsonrpc": "2.0",
"id": json_req.id,
"error": {
"code": -32000,
"message": "Server not initialized. Call 'initialize' first."
}
}))
.into_response();
}
}
let response = match json_req.method.as_str() {
"initialize" => {
let (result, session_id) = handle_initialize(server_instance);
let resp = JsonRpcResponse::success(id.clone(), result);
let body = Json(serde_json::to_value(&resp).unwrap()).into_response();
if let Some(sid) = session_id {
let mut response = body;
let sid_header = HeaderValue::from_str(&sid)
.unwrap_or_else(|_| HeaderValue::from_static("unknown"));
response.headers_mut().insert("Mcp-Session-Id", sid_header);
return response;
}
return body;
}
"ping" => Ok(handle_ping()),
"tools/call" => {
let _guard = incoming_session_id
.as_ref()
.map(|sid| McpServer::in_flight_guard(server_instance, sid));
let advisory = incoming_session_id
.as_deref()
.map(|sid| (server_instance.as_ref(), sid));
handle_tool_call_timed(&state, handlers, &json_req, transport_started, advisory).await
}
"tools/list" => Ok(list_tools_json(handlers)),
"prompts/list" => Ok(list_prompts_json()),
"prompts/get" => handle_prompt_get(&json_req),
"resources/list" => Ok(list_resources_json()),
"resources/read" => handle_resource_read(&json_req),
_ => Err(JsonRpcError::method_not_found(json_req.method.clone())),
};
let resp = match response {
Ok(result) => {
debug!("Request completed successfully");
JsonRpcResponse::success(id, result)
}
Err(e) => {
warn!("Request failed: {}", e);
JsonRpcResponse::error(id, e)
}
};
Json(serde_json::to_value(&resp).unwrap()).into_response()
}
pub async fn handle_tool_call(
registry: &Arc<ProjectRegistry>,
handlers: &[ToolHandler],
req: &JsonRpcRequest,
) -> Result<Value, JsonRpcError> {
handle_tool_call_timed(registry, handlers, req, Instant::now(), None).await
}
async fn handle_tool_call_timed(
registry: &Arc<ProjectRegistry>,
handlers: &[ToolHandler],
req: &JsonRpcRequest,
transport_started: Instant,
advisory: Option<(&McpServer, &str)>,
) -> Result<Value, JsonRpcError> {
let handler_started = Instant::now();
let tool_call = req.extract_tool_call()?;
debug!("Tool call: name={}", tool_call.name);
let handler = handlers
.iter()
.find(|h| h.name() == tool_call.name)
.ok_or_else(|| JsonRpcError::method_not_found(tool_call.name.clone()))?;
let call_args = tool_call.arguments.clone();
let call_name = tool_call.name.clone();
let (handler_result, mut timings) =
collect_request_timings(handler.execute(registry, tool_call.arguments)).await;
let handler_ms = elapsed_ms(handler_started);
timings.handler_ms = handler_ms;
timings.transport_queue_ms = handler_started
.saturating_duration_since(transport_started)
.as_millis()
.min(u64::MAX as u128) as u64;
timings.total_ms = elapsed_ms(transport_started);
debug!(
tool = %call_name,
handler_ms,
transport_queue_ms = timings.transport_queue_ms,
total_ms = timings.total_ms,
"MCP tool call complete"
);
match handler_result {
Ok(mut value) => {
if let Some((server, session_id)) = advisory {
server.apply_freshness_advisory(
session_id,
call_args.get("project_path").and_then(Value::as_str),
&mut value,
);
}
let trimmed = crate::cli::mcp::output::trim_llm_payload(&call_name, &value);
let rendered =
crate::cli::mcp::output::render_tool_output_plain(&call_name, &trimmed, &call_args);
let is_substantive = rendered.trim().lines().count() > 1 || rendered.trim().len() > 80;
let payload = if is_substantive {
rendered
} else {
serde_json::to_string_pretty(&trimmed)
.unwrap_or_else(|_| "Error serializing result".to_string())
};
Ok(serde_json::json!({
"content": [
{
"type": "text",
"text": payload
}
],
"isError": false,
"_meta": { "timings": timings }
}))
}
Err(e) => {
warn!("Tool execution failed: {}", e);
Ok(serde_json::json!({
"content": [
{
"type": "text",
"text": format!("Error: {}", e)
}
],
"isError": true,
"_meta": { "timings": timings }
}))
}
}
}
pub fn list_tools_json(handlers: &[ToolHandler]) -> Value {
let tools: Vec<_> = handlers
.iter()
.map(|handler| {
serde_json::json!({
"name": handler.name(),
"description": handler.description(),
"inputSchema": handler.argument_schema()
})
})
.collect();
serde_json::json!({ "tools": tools })
}
async fn list_tools_handler(headers: HeaderMap) -> Json<Value> {
if let Some(sid) = headers.get("Mcp-Session-Id").and_then(|v| v.to_str().ok()) {
if let Some(server) = SERVER_INSTANCE.get() {
if let Some(mut entry) = server.session_handshakes.get_mut(sid) {
entry.1 = Instant::now();
if !entry.0 {
return Json(serde_json::json!({
"error": "Invalid session. Call 'initialize' first."
}));
}
}
}
}
if SERVER_INSTANCE.get().is_none() {
return Json(serde_json::json!({
"error": "Server instance not initialized"
}));
}
let handlers = match HANDLERS.get() {
Some(h) => h,
None => {
return Json(serde_json::json!({
"error": "Handlers not initialized"
}));
}
};
Json(list_tools_json(handlers))
}
async fn health_check_handler() -> Json<Value> {
Json(serde_json::json!({
"status": "ok",
"service": "leindex-mcp-server",
"version": env!("CARGO_PKG_VERSION")
}))
}
pub struct InFlightGuard {
server: Arc<McpServer>,
session_id: Arc<str>,
}
impl Drop for InFlightGuard {
fn drop(&mut self) {
self.server.end_request(&self.session_id);
}
}
#[cfg(unix)]
pub struct SocketCleanupGuard {
path: std::path::PathBuf,
}
#[cfg(unix)]
impl Drop for SocketCleanupGuard {
fn drop(&mut self) {
if self.path.exists() {
let _ = std::fs::remove_file(&self.path);
debug!("Cleaned up socket file: {}", self.path.display());
}
}
}
#[cfg(unix)]
impl McpServer {
pub async fn run_socket(&self, socket_path: &std::path::Path) -> anyhow::Result<()> {
use tokio::net::UnixListener;
if socket_path.exists() {
std::fs::remove_file(socket_path).context("Failed to remove existing socket file")?;
}
if let Some(parent) = socket_path.parent() {
std::fs::create_dir_all(parent).context("Failed to create socket directory")?;
}
let listener = UnixListener::bind(socket_path)
.with_context(|| format!("Failed to bind Unix socket at {}", socket_path.display()))?;
let _guard = SocketCleanupGuard {
path: socket_path.to_path_buf(),
};
info!(
"MCP server listening on Unix socket: {}",
socket_path.display()
);
loop {
let (stream, _addr) = listener
.accept()
.await
.context("Failed to accept connection")?;
let session_id = generate_session_id();
self.session_handshakes.insert(
Arc::<str>::from(session_id.as_str()),
(false, Instant::now()),
);
tokio::spawn(handle_socket_connection(
stream,
session_id,
self.session_handshakes.clone(),
self.handshake_complete.clone(),
));
}
#[allow(unreachable_code)]
{
Ok(())
}
}
}
#[cfg(unix)]
#[cfg(unix)]
async fn read_bounded_line<R: tokio::io::AsyncBufRead + Unpin>(
reader: &mut R,
max: usize,
) -> std::io::Result<Option<String>> {
use tokio::io::AsyncBufReadExt;
let mut out: Vec<u8> = Vec::new();
loop {
let buf = reader.fill_buf().await?;
if buf.is_empty() {
return Ok(if out.is_empty() {
None
} else {
Some(String::from_utf8_lossy(&out).into_owned())
});
}
let nl = buf.iter().position(|&b| b == b'\n');
let take = nl.map(|i| i + 1).unwrap_or(buf.len());
if out.len().saturating_add(take) > max {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"line exceeds max length",
));
}
out.extend_from_slice(&buf[..take]);
reader.consume(take);
if nl.is_some() {
return Ok(Some(String::from_utf8_lossy(&out).into_owned()));
}
}
}
#[cfg(unix)]
#[derive(Debug)]
enum SocketFrame {
Message {
payload: String,
content_length: bool,
},
Error {
response: String,
content_length: bool,
},
}
#[cfg(unix)]
async fn read_socket_frame<R>(
reader: &mut R,
session_id: &str,
first_frame: bool,
) -> Option<SocketFrame>
where
R: tokio::io::AsyncBufRead + tokio::io::AsyncRead + Unpin,
{
use tokio::io::AsyncReadExt;
const MAX_LINE_LENGTH: usize = 10_240;
const MAX_PAYLOAD_SIZE: usize = 10_485_760;
let line_timeout = socket_read_timeout(first_frame);
let line = match tokio::time::timeout(line_timeout, read_bounded_line(reader, MAX_PAYLOAD_SIZE))
.await
{
Ok(Ok(Some(line))) => line,
Ok(Ok(None)) => return None,
Ok(Err(error)) => {
debug!(
"Socket read error / line too long (session {}): {}",
session_id, error
);
let error_response = JsonRpcResponse::error(
serde_json::Value::Null,
JsonRpcError::new(-32600, "request payload exceeds maximum size"),
);
return serde_json::to_string(&error_response).ok().map(|response| {
SocketFrame::Error {
response,
content_length: false,
}
});
}
Err(_) => {
debug!("Socket read timed out (session {})", session_id);
return None;
}
};
let line_trim = line.trim_end();
if line_trim.is_empty() {
return Some(SocketFrame::Message {
payload: String::new(),
content_length: false,
});
}
if !line_trim
.to_ascii_lowercase()
.starts_with("content-length:")
{
return Some(SocketFrame::Message {
payload: line_trim.to_string(),
content_length: false,
});
}
let len_str = line_trim.split(':').nth(1).unwrap_or("").trim();
let length = match len_str.parse::<usize>() {
Ok(length) => length,
Err(error) => {
debug!("Invalid Content-Length header: {}", error);
let response = JsonRpcResponse::error(
serde_json::Value::Null,
JsonRpcError::new(-32600, "invalid Content-Length header"),
);
return serde_json::to_string(&response)
.ok()
.map(|response| SocketFrame::Error {
response,
content_length: false,
});
}
};
if length > MAX_PAYLOAD_SIZE {
debug!(
"Payload too large (session {}): {} bytes",
session_id, length
);
let response = JsonRpcResponse::error(
serde_json::Value::Null,
JsonRpcError::new(-32600, "request payload exceeds maximum size"),
);
return serde_json::to_string(&response)
.ok()
.map(|response| SocketFrame::Error {
response,
content_length: true,
});
}
loop {
let header = match tokio::time::timeout(
SOCKET_READ_TIMEOUT,
read_bounded_line(reader, MAX_LINE_LENGTH),
)
.await
{
Ok(Ok(Some(header))) => header,
Ok(Ok(None)) | Ok(Err(_)) | Err(_) => return None,
};
if header.trim().is_empty() {
break;
}
}
let mut buffer = vec![0u8; length];
match tokio::time::timeout(SOCKET_READ_TIMEOUT, reader.read_exact(&mut buffer)).await {
Ok(Ok(_)) => Some(SocketFrame::Message {
payload: String::from_utf8_lossy(&buffer).into_owned(),
content_length: true,
}),
Ok(Err(error)) => {
debug!("Failed to read JSON payload: {}", error);
None
}
Err(_) => {
debug!("Socket payload read timed out (session {})", session_id);
None
}
}
}
#[cfg(unix)]
async fn write_socket_frame<W>(writer: &mut W, response: &str, content_length: bool) -> bool
where
W: tokio::io::AsyncWrite + Unpin,
{
use tokio::io::AsyncWriteExt;
let message = if content_length {
format!("Content-Length: {}\r\n\r\n{}", response.len(), response)
} else {
format!("{}\n", response)
};
writer.write_all(message.as_bytes()).await.is_ok() && writer.flush().await.is_ok()
}
#[cfg(unix)]
async fn handle_socket_connection(
stream: tokio::net::UnixStream,
session_id: String,
session_handshakes: Arc<DashMap<Arc<str>, (bool, Instant)>>,
handshake_complete: Arc<AtomicBool>,
) {
use tokio::io::BufReader;
debug!("Accepted Unix socket connection (session: {})", session_id);
let (reader, mut writer) = stream.into_split();
let mut reader = BufReader::new(reader);
let mut first_frame = true;
loop {
let Some(frame) = read_socket_frame(&mut reader, &session_id, first_frame).await else {
break;
};
first_frame = false;
let (json_payload, content_length) = match frame {
SocketFrame::Message {
payload,
content_length,
} => (payload, content_length),
SocketFrame::Error {
response,
content_length,
} => {
let _ = write_socket_frame(&mut writer, &response, content_length).await;
break;
}
};
if json_payload.is_empty() {
continue;
}
let Some(response) = handle_socket_message(
&json_payload,
&session_id,
&session_handshakes,
&handshake_complete,
)
.await
else {
continue;
};
if !write_socket_frame(&mut writer, &response, content_length).await {
break;
}
}
session_handshakes.remove(session_id.as_str());
debug!("Socket connection closed (session: {})", session_id);
}
#[cfg(unix)]
async fn handle_socket_message(
json_payload: &str,
session_id: &str,
session_handshakes: &Arc<DashMap<Arc<str>, (bool, Instant)>>,
handshake_complete: &Arc<AtomicBool>,
) -> Option<String> {
use super::protocol::{JsonRpcMessage, JsonRpcResponse};
use crate::cli::mcp::server::{HANDLERS, SERVER_STATE, list_tools_json};
let transport_started = Instant::now();
let message = match JsonRpcMessage::from_json(json_payload) {
Ok(m) => m,
Err(e) => {
let error_response = JsonRpcResponse::error(serde_json::Value::Null, e);
return serde_json::to_string(&error_response).ok();
}
};
match message {
JsonRpcMessage::Notification(notification) => {
debug!("Ignoring notification on socket: {}", notification.method);
None
}
JsonRpcMessage::Request(request) => {
let request_id = request.id.clone().unwrap_or(serde_json::Value::Null);
let method_name = request.method.clone();
if request.id.is_none() {
debug!("Ignoring notification: {}", method_name);
return None;
}
let state = match SERVER_STATE.get() {
Some(s) => s,
None => {
let resp = JsonRpcResponse::error(
request_id,
super::protocol::JsonRpcError::new(-32603, "Server state not initialized"),
);
return serde_json::to_string(&resp).ok();
}
};
let handlers = match HANDLERS.get() {
Some(h) => h,
None => {
let resp = JsonRpcResponse::error(
request_id,
super::protocol::JsonRpcError::new(-32603, "Handlers not initialized"),
);
return serde_json::to_string(&resp).ok();
}
};
if method_name != "initialize" && method_name != "ping" {
let handshaked = match session_handshakes.get_mut(session_id) {
Some(mut entry) => {
entry.1 = Instant::now();
entry.0
}
_ => false,
};
if !handshaked {
let resp = JsonRpcResponse::error(
request_id,
super::protocol::JsonRpcError::new(
-32600,
"Server not initialized. Call 'initialize' first.",
),
);
return serde_json::to_string(&resp).ok();
}
}
let response = match method_name.as_str() {
"initialize" => {
handshake_complete.store(true, Ordering::SeqCst);
session_handshakes.insert(Arc::<str>::from(session_id), (true, Instant::now()));
let result = serde_json::json!({
"protocolVersion": "2024-11-05",
"capabilities": {
"tools": { "listChanged": true },
"prompts": { "listChanged": true },
"resources": { "listChanged": true, "subscribe": false },
"logging": {},
"progress": true
},
"serverInfo": {
"name": "leindex",
"version": env!("CARGO_PKG_VERSION"),
"description": "LeIndex MCP Server - Semantic code indexing and analysis with PDG-based tools"
}
});
JsonRpcResponse::success(request_id, result)
}
"ping" => JsonRpcResponse::success(request_id, serde_json::json!({})),
"tools/call" => {
let advisory = SERVER_INSTANCE
.get()
.map(|server| (server.as_ref(), session_id));
let result = handle_tool_call_timed(
state,
handlers,
&request,
transport_started,
advisory,
)
.await;
JsonRpcResponse::from_result(request_id, result)
}
"tools/list" => JsonRpcResponse::success(request_id, list_tools_json(handlers)),
"prompts/list" => JsonRpcResponse::success(request_id, list_prompts_json()),
"prompts/get" => {
let result = handle_prompt_get(&request);
match result {
Ok(value) => JsonRpcResponse::success(request_id, value),
Err(e) => JsonRpcResponse::error(request_id, e),
}
}
"resources/list" => JsonRpcResponse::success(request_id, list_resources_json()),
"resources/read" => {
let result = handle_resource_read(&request);
match result {
Ok(value) => JsonRpcResponse::success(request_id, value),
Err(e) => JsonRpcResponse::error(request_id, e),
}
}
_ => JsonRpcResponse::error(
request_id,
super::protocol::JsonRpcError::method_not_found(method_name),
),
};
serde_json::to_string(&response).ok()
}
}
}
#[cfg(test)]
#[path = "server_test.rs"]
mod tests;
#[path = "prompts_resources.rs"]
mod prompts_resources;
pub use prompts_resources::{
Prompt, PromptArgument, PromptContent, PromptMessage, Resource, ResourceContent, get_prompt,
get_prompts, get_resource, get_resources, handle_prompt_get, handle_resource_read,
list_prompts_json, list_resources_json,
};