use super::*;
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn test_server_config_default() {
let config = McpServerConfig::default();
assert_eq!(
config.bind_address,
SocketAddr::from(([127, 0, 0, 1], DEFAULT_MCP_PORT))
);
assert!(config.bind_address.ip().is_loopback());
assert!(config.bind_address.port() >= 10000);
}
#[cfg(unix)]
#[tokio::test]
async fn test_read_bounded_line_rejects_oversized_input_and_handles_eof() {
use tokio::io::BufReader;
let mut oversized = BufReader::new(std::io::Cursor::new(b"12345".to_vec()));
let error = read_bounded_line(&mut oversized, 4).await.unwrap_err();
assert_eq!(error.kind(), std::io::ErrorKind::InvalidData);
let mut exact = BufReader::new(std::io::Cursor::new(b"1234\n".to_vec()));
assert_eq!(
read_bounded_line(&mut exact, 5).await.unwrap().as_deref(),
Some("1234\n")
);
let mut valid = BufReader::new(std::io::Cursor::new(b"ok\n".to_vec()));
assert_eq!(
read_bounded_line(&mut valid, 3).await.unwrap().as_deref(),
Some("ok\n")
);
let mut eof = BufReader::new(std::io::Cursor::new(Vec::<u8>::new()));
assert!(read_bounded_line(&mut eof, 4).await.unwrap().is_none());
}
#[cfg(unix)]
#[tokio::test]
async fn test_socket_connection_preserves_newline_and_content_length_framing() {
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
let (mut client, server) = tokio::net::UnixStream::pair().unwrap();
let session_handshakes = Arc::new(DashMap::new());
let handshake_complete = Arc::new(AtomicBool::new(false));
let task = tokio::spawn(handle_socket_connection(
server,
"framing-test".to_string(),
session_handshakes.clone(),
handshake_complete,
));
client.write_all(b"not-json\n").await.unwrap();
let (reader, mut writer) = client.into_split();
let mut reader = BufReader::new(reader);
let mut newline_response = String::new();
reader.read_line(&mut newline_response).await.unwrap();
assert!(newline_response.contains("-32700") || newline_response.contains("-32600"));
let payload = b"not-json";
writer
.write_all(format!("Content-Length: {}\r\n\r\n", payload.len()).as_bytes())
.await
.unwrap();
writer.write_all(payload).await.unwrap();
writer.flush().await.unwrap();
let mut header = String::new();
reader.read_line(&mut header).await.unwrap();
assert!(header.to_ascii_lowercase().starts_with("content-length:"));
let length = header
.split_once(':')
.and_then(|(_, value)| value.trim().parse::<usize>().ok())
.expect("response content length");
let mut blank = String::new();
reader.read_line(&mut blank).await.unwrap();
assert!(blank.trim().is_empty());
let mut response = vec![0; length];
reader.read_exact(&mut response).await.unwrap();
let response = String::from_utf8(response).unwrap();
assert!(response.contains("-32700") || response.contains("-32600"));
writer.shutdown().await.unwrap();
task.await.unwrap();
assert!(!session_handshakes.contains_key("framing-test"));
}
#[cfg(unix)]
#[tokio::test]
async fn test_socket_connection_closes_on_malformed_or_incomplete_headers() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (mut client, server) = tokio::net::UnixStream::pair().unwrap();
let task = tokio::spawn(handle_socket_connection(
server,
"malformed-header-test".to_string(),
Arc::new(DashMap::new()),
Arc::new(AtomicBool::new(false)),
));
client.write_all(b"Content-Length: nope\r\n").await.unwrap();
let mut response = Vec::new();
client.read_to_end(&mut response).await.unwrap();
assert!(String::from_utf8_lossy(&response).contains("-32600"));
task.await.unwrap();
let (mut client, server) = tokio::net::UnixStream::pair().unwrap();
let task = tokio::spawn(handle_socket_connection(
server,
"incomplete-header-test".to_string(),
Arc::new(DashMap::new()),
Arc::new(AtomicBool::new(false)),
));
client
.write_all(b"Content-Length: 8\r\n\r\nshort")
.await
.unwrap();
client.shutdown().await.unwrap();
let mut response = Vec::new();
client.read_to_end(&mut response).await.unwrap();
assert!(response.is_empty());
tokio::time::timeout(std::time::Duration::from_secs(1), task)
.await
.expect("incomplete payload must close promptly")
.unwrap();
}
#[cfg(unix)]
#[test]
fn test_socket_cleanup_guard_removes_file() {
let dir = std::env::temp_dir().join("leindex_test_socket_guard");
std::fs::create_dir_all(&dir).unwrap();
let socket_path = dir.join("test.sock");
std::fs::write(&socket_path, b"").unwrap();
assert!(socket_path.exists());
{
let _guard = SocketCleanupGuard {
path: socket_path.clone(),
};
}
assert!(!socket_path.exists());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn test_concurrent_session_handshake_isolation() {
let registry = Arc::new(ProjectRegistry::new(5));
let server = McpServer {
config: McpServerConfig::default(),
_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()),
};
let (result1, sid1) = handle_initialize_for_test(&server);
let (result2, sid2) = handle_initialize_for_test(&server);
let (result3, sid3) = handle_initialize_for_test(&server);
assert!(result1.get("protocolVersion").is_some());
assert!(result2.get("protocolVersion").is_some());
assert!(result3.get("protocolVersion").is_some());
assert_ne!(sid1, sid2);
assert_ne!(sid2, sid3);
assert_ne!(sid1, sid3);
assert_eq!(server.active_session_count(), 3);
{
assert!(server.session_handshakes.get(sid1.as_str()).unwrap().0);
assert!(server.session_handshakes.get(sid2.as_str()).unwrap().0);
assert!(server.session_handshakes.get(sid3.as_str()).unwrap().0);
}
}
#[test]
fn test_session_isolation_per_session() {
let registry = Arc::new(ProjectRegistry::new(5));
let server = McpServer {
config: McpServerConfig::default(),
_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()),
};
let (_, sid1) = handle_initialize_for_test(&server);
let (_, sid2) = handle_initialize_for_test(&server);
server.session_handshakes.remove(sid1.as_str());
assert!(server.session_handshakes.get(sid2.as_str()).is_some());
assert!(server.session_handshakes.get(sid2.as_str()).unwrap().0);
assert_eq!(server.active_session_count(), 1);
}
#[test]
fn test_freshness_advisory_is_once_per_session_and_generation() {
let registry = Arc::new(ProjectRegistry::new(5));
let server = McpServer {
config: McpServerConfig::default(),
_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()),
};
let (_, session_id) = handle_initialize_for_test(&server);
let response = || {
serde_json::json!({
"project_path": "/tmp/project",
"_meta": {"freshness": {
"generation": 7,
"warning": "refresh recommended"
}}
})
};
let mut first = response();
server.apply_freshness_advisory(&session_id, None, &mut first);
assert_eq!(
first["_meta"]["freshness"]["advisory"],
"refresh recommended"
);
assert!(first["_meta"]["freshness"].get("warning").is_none());
let mut second = response();
server.apply_freshness_advisory(&session_id, None, &mut second);
assert!(second["_meta"]["freshness"]["advisory"].is_null());
let mut new_generation = response();
new_generation["_meta"]["freshness"]["generation"] = serde_json::json!(8);
server.apply_freshness_advisory(&session_id, None, &mut new_generation);
assert_eq!(
new_generation["_meta"]["freshness"]["advisory"],
"refresh recommended"
);
}
#[test]
fn test_stale_session_cleanup() {
let registry = Arc::new(ProjectRegistry::new(5));
let server = McpServer {
config: McpServerConfig::default(),
_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()),
};
let (_, sid) = handle_initialize_for_test(&server);
assert_eq!(server.active_session_count(), 1);
if let Some(mut entry) = server.session_handshakes.get_mut(sid.as_str()) {
entry.1 = Instant::now() - std::time::Duration::from_secs(600);
}
let removed = server.cleanup_stale_sessions(std::time::Duration::from_secs(60));
assert_eq!(removed, 1);
assert_eq!(server.active_session_count(), 0);
}
fn handle_initialize_for_test(server: &McpServer) -> (Value, String) {
let (result, sid) = handle_initialize(server);
(result, sid.unwrap())
}
#[test]
fn test_handle_initialize_does_not_evict_in_flight_session() {
let _env_guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let registry = Arc::new(ProjectRegistry::new(5));
let server = McpServer {
config: McpServerConfig::default(),
_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()),
};
let now = Instant::now();
for i in 0..DEFAULT_MAX_HTTP_SESSIONS {
let sid = format!("sess-{i:04}");
server
.session_handshakes
.insert(Arc::<str>::from(sid.as_str()), (true, now));
}
assert_eq!(server.active_session_count(), DEFAULT_MAX_HTTP_SESSIONS);
let in_flight_sid = "sess-0000".to_string();
server.begin_request(&in_flight_sid);
assert!(server.session_in_flight(&in_flight_sid));
let (_, new_sid) = handle_initialize_for_test(&server);
assert!(
server
.session_handshakes
.contains_key(in_flight_sid.as_str()),
"in_flight session {} was evicted during initialize",
in_flight_sid,
);
assert!(
server.session_handshakes.contains_key(new_sid.as_str()),
"newly initialized session {} was not registered",
new_sid,
);
assert!(
server.active_session_count() <= DEFAULT_MAX_HTTP_SESSIONS + 1,
"session count {} exceeds cap + 1",
server.active_session_count(),
);
server.end_request(&in_flight_sid);
}
#[test]
fn test_max_http_sessions_env_override() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
unsafe {
std::env::set_var(MAX_SESSIONS_ENV, "42");
}
assert_eq!(max_http_sessions(), 42);
unsafe {
std::env::remove_var(MAX_SESSIONS_ENV);
}
assert_eq!(max_http_sessions(), DEFAULT_MAX_HTTP_SESSIONS);
unsafe {
std::env::set_var(MAX_SESSIONS_ENV, "not-a-number");
}
assert_eq!(max_http_sessions(), DEFAULT_MAX_HTTP_SESSIONS);
unsafe {
std::env::remove_var(MAX_SESSIONS_ENV);
}
}
#[test]
fn test_begin_request_clones_cached_arc_str_key() {
let registry = Arc::new(ProjectRegistry::new(5));
let server = McpServer {
config: McpServerConfig::default(),
_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()),
};
let sid = Arc::<str>::from("sess-arc-key");
let cached: Arc<str> = sid.clone();
server
.session_handshakes
.insert(sid, (true, Instant::now()));
server.begin_request("sess-arc-key");
let stored = server
.in_flight
.iter()
.next()
.expect("expected one in_flight entry");
assert_eq!(
stored.key().as_ptr(),
cached.as_ptr(),
"in_flight key should be the same Arc<str> allocation as session_handshakes"
);
assert_eq!(stored.key().as_ref(), "sess-arc-key");
}
#[test]
fn test_begin_request_allocates_when_session_unknown() {
let registry = Arc::new(ProjectRegistry::new(5));
let server = McpServer {
config: McpServerConfig::default(),
_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.begin_request("unregistered-session");
assert!(
server.session_in_flight("unregistered-session"),
"begin_request must insert even when session is unregistered"
);
let stored = server
.in_flight
.get("unregistered-session")
.expect("lookup by &str must work");
assert_eq!(stored.key().as_ref(), "unregistered-session");
}
#[tokio::test]
async fn test_bind_with_fallback_breaks_on_port_overflow() {
let high = SocketAddr::from(([127, 0, 0, 1], u16::MAX));
let _occupying = tokio::net::TcpListener::bind(high).await.unwrap();
let preferred = SocketAddr::from(([127, 0, 0, 1], u16::MAX - 5));
let result = tokio::time::timeout(
std::time::Duration::from_secs(5),
bind_with_fallback(preferred),
)
.await;
let listener = result
.expect("bind_with_fallback must terminate within 5s")
.expect("ephemeral fallback must succeed on a free IP");
assert_eq!(
listener.local_addr().unwrap().ip(),
preferred.ip(),
"ephemeral fallback must use the preferred IP"
);
}
#[tokio::test]
async fn test_bind_with_fallback_uses_preferred_when_free() {
let preferred = SocketAddr::from(([127, 0, 0, 1], 0));
let listener = bind_with_fallback(preferred).await.unwrap();
let bound = listener.local_addr().unwrap();
assert_eq!(bound.ip(), preferred.ip());
}
#[tokio::test]
async fn test_bind_with_fallback_port_zero_skips_fallback_loop() {
let preferred = SocketAddr::from(([127, 0, 0, 1], 0));
let listener = bind_with_fallback(preferred).await.unwrap();
let bound = listener.local_addr().unwrap();
assert_ne!(
bound.port(),
0,
"OS-assigned ephemeral port must be non-zero; got port 0 (bind did not actually request ephemeral)"
);
assert_eq!(bound.ip(), preferred.ip());
}
#[cfg(unix)]
#[test]
fn test_socket_timeout_contract_distinguishes_initial_and_subsequent_frames() {
assert_eq!(
socket_read_timeout(true),
std::time::Duration::from_secs(120)
);
assert_eq!(
socket_read_timeout(false),
std::time::Duration::from_secs(30)
);
assert!(socket_read_timeout(true) > socket_read_timeout(false));
}