use crate::tools::native::{NativeTool, ParameterSchema, ToolInput, ToolOutput, ToolSchema};
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::time::{Duration, Instant};
const MAX_TOTAL_REQUESTS: usize = 100_000;
pub struct LoadTestingTool;
impl NativeTool for LoadTestingTool {
fn name(&self) -> &str {
"load_testing"
}
fn description(&self) -> &str {
"Generate HTTP load against a target URL with concurrent users and measure performance metrics."
}
fn schema(&self) -> ToolSchema {
ToolSchema {
name: self.name().to_owned(),
description: self.description().to_owned(),
parameters: vec![
ParameterSchema {
name: "target_url".to_owned(),
description: "URL to load test".to_owned(),
param_type: "string".to_owned(),
required: true,
},
ParameterSchema {
name: "concurrent_users".to_owned(),
description: "Number of concurrent users (default: 10)".to_owned(),
param_type: "number".to_owned(),
required: false,
},
ParameterSchema {
name: "duration_seconds".to_owned(),
description: "Test duration in seconds (default: 10)".to_owned(),
param_type: "number".to_owned(),
required: false,
},
],
}
}
fn execute(&self, input: ToolInput) -> Pin<Box<dyn Future<Output = ToolOutput> + Send + '_>> {
Box::pin(async move {
let target_url = match input.get_str("target_url") {
Some(url) => url.to_string(),
None => return ToolOutput::err("missing required parameter: target_url"),
};
if !crate::server::ssrf::is_safe_url(&target_url) {
return ToolOutput::err(
"target_url rejected: cannot target private/internal networks",
);
}
let concurrent_users = input.get_u64("concurrent_users").unwrap_or(10) as usize;
let duration_secs = input.get_u64("duration_seconds").unwrap_or(10);
let concurrent_users = concurrent_users.min(500);
let duration = Duration::from_secs(duration_secs.clamp(1, 300));
tracing::info!(
target_url = %target_url,
concurrent_users,
duration_secs = duration.as_secs(),
"load test starting"
);
match run_load_test(&target_url, concurrent_users, duration).await {
Ok(result) => {
tracing::info!(
total_requests = result.total_requests,
throughput_rps = result.throughput_rps,
p99_latency_ms = result.p99_latency_ms,
"load test completed"
);
ToolOutput::ok(serde_json::to_value(result).unwrap_or_default())
}
Err(e) => ToolOutput::err(format!("load test failed: {e}")),
}
})
}
}
#[derive(serde::Serialize)]
struct LoadTestResult {
target_url: String,
concurrent_users: usize,
duration_seconds: u64,
total_requests: u64,
successful_requests: u64,
failed_requests: u64,
avg_latency_ms: f64,
p50_latency_ms: f64,
p95_latency_ms: f64,
p99_latency_ms: f64,
min_latency_ms: f64,
max_latency_ms: f64,
throughput_rps: f64,
error_rate: f64,
status_codes: HashMap<u16, u64>,
}
async fn run_load_test(
target_url: &str,
concurrent_users: usize,
duration: Duration,
) -> Result<LoadTestResult, String> {
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.pool_max_idle_per_host(concurrent_users)
.build()
.map_err(|e| e.to_string())?;
let url = target_url.to_string();
let deadline = Instant::now() + duration;
let mut handles = Vec::with_capacity(concurrent_users);
let max_requests_per_user = MAX_TOTAL_REQUESTS / concurrent_users.max(1);
for _ in 0..concurrent_users {
let client = client.clone();
let url = url.clone();
handles.push(tokio::spawn(async move {
let mut latencies = Vec::new();
let mut status_counts: HashMap<u16, u64> = HashMap::new();
let mut errors: u64 = 0;
while Instant::now() < deadline
&& latencies.len() + (errors as usize) < max_requests_per_user
{
let start = Instant::now();
match client.get(&url).send().await {
Ok(resp) => {
let status = resp.status().as_u16();
*status_counts.entry(status).or_default() += 1;
latencies.push(start.elapsed());
}
Err(_) => {
errors += 1;
latencies.push(start.elapsed());
}
}
}
(latencies, status_counts, errors)
}));
}
let mut all_latencies: Vec<Duration> = Vec::new();
let mut all_status_codes: HashMap<u16, u64> = HashMap::new();
let mut total_errors: u64 = 0;
for handle in handles {
match handle.await {
Ok((latencies, statuses, errors)) => {
all_latencies.extend(latencies);
for (code, count) in statuses {
*all_status_codes.entry(code).or_default() += count;
}
total_errors += errors;
}
Err(_) => total_errors += 1,
}
}
let total_requests = all_latencies.len() as u64;
let successful_requests = total_requests.saturating_sub(total_errors);
if all_latencies.is_empty() {
return Err("no requests completed".to_string());
}
all_latencies.sort();
let to_ms = |d: Duration| d.as_secs_f64() * 1000.0;
let avg_latency =
all_latencies.iter().map(|d| to_ms(*d)).sum::<f64>() / all_latencies.len() as f64;
let last = all_latencies.len() - 1;
let p50 = to_ms(all_latencies[(all_latencies.len() * 50 / 100).min(last)]);
let p95 = to_ms(all_latencies[(all_latencies.len() * 95 / 100).min(last)]);
let p99 = to_ms(all_latencies[(all_latencies.len() * 99 / 100).min(last)]);
let min_latency = to_ms(all_latencies[0]);
let max_latency = to_ms(*all_latencies.last().unwrap());
let elapsed = duration.as_secs_f64();
let throughput = total_requests as f64 / elapsed;
let error_rate = if total_requests > 0 {
total_errors as f64 / total_requests as f64
} else {
0.0
};
Ok(LoadTestResult {
target_url: url,
concurrent_users,
duration_seconds: duration.as_secs(),
total_requests,
successful_requests,
failed_requests: total_errors,
avg_latency_ms: (avg_latency * 100.0).round() / 100.0,
p50_latency_ms: (p50 * 100.0).round() / 100.0,
p95_latency_ms: (p95 * 100.0).round() / 100.0,
p99_latency_ms: (p99 * 100.0).round() / 100.0,
min_latency_ms: (min_latency * 100.0).round() / 100.0,
max_latency_ms: (max_latency * 100.0).round() / 100.0,
throughput_rps: (throughput * 100.0).round() / 100.0,
error_rate: (error_rate * 10000.0).round() / 10000.0,
status_codes: all_status_codes,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn load_testing_name() {
assert_eq!(LoadTestingTool.name(), "load_testing");
}
#[test]
fn load_testing_schema_has_target_url() {
let schema = LoadTestingTool.schema();
assert!(
schema
.parameters
.iter()
.any(|p| p.name == "target_url" && p.required)
);
}
#[tokio::test]
async fn load_testing_missing_url() {
let tool = LoadTestingTool;
let input = ToolInput {
parameters: HashMap::new(),
};
let output = tool.execute(input).await;
assert!(!output.success);
assert!(output.error.unwrap().contains("target_url"));
}
async fn mock_server() -> (String, tokio::task::JoinHandle<()>) {
use axum::{Router, routing::get};
let app = Router::new().route("/ok", get(|| async { "OK" }));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let handle = tokio::spawn(async move {
axum::serve(listener, app).await.ok();
});
(format!("http://{addr}"), handle)
}
#[tokio::test]
async fn load_test_ssrf_rejects_private_ip() {
let tool = LoadTestingTool;
let mut params = HashMap::new();
params.insert(
"target_url".to_owned(),
serde_json::Value::String("http://10.0.0.1/secret".to_owned()),
);
let output = tool.execute(ToolInput { parameters: params }).await;
assert!(!output.success);
assert!(output.error.unwrap().contains("private"));
}
#[tokio::test]
async fn run_load_test_against_mock_server() {
let (base_url, handle) = mock_server().await;
let url = format!("{base_url}/ok");
let result = run_load_test(&url, 2, Duration::from_secs(1))
.await
.expect("load test should succeed");
assert!(result.total_requests > 0);
assert!(result.successful_requests > 0);
assert_eq!(result.failed_requests, 0);
assert!(result.avg_latency_ms > 0.0);
assert!(result.p50_latency_ms > 0.0);
assert!(result.throughput_rps > 0.0);
assert!(result.error_rate < f64::EPSILON);
assert_eq!(
*result.status_codes.get(&200).unwrap_or(&0),
result.total_requests
);
handle.abort();
}
#[tokio::test]
async fn run_load_test_records_error_status_codes() {
use axum::{Router, http::StatusCode, routing::get};
let app = Router::new().route("/bad", get(|| async { StatusCode::INTERNAL_SERVER_ERROR }));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let handle = tokio::spawn(async move {
axum::serve(listener, app).await.ok();
});
let url = format!("http://{addr}/bad");
let result = run_load_test(&url, 1, Duration::from_secs(1))
.await
.expect("load test should complete");
assert!(result.total_requests > 0);
assert!(result.status_codes.contains_key(&500));
handle.abort();
}
#[tokio::test]
async fn run_load_test_connection_refused() {
let result = run_load_test("http://127.0.0.1:1/nope", 1, Duration::from_secs(1)).await;
if let Ok(r) = result {
assert!(r.failed_requests > 0 || r.error_rate > 0.0);
}
}
#[tokio::test]
async fn load_test_execute_trait_ssrf_blocks_localhost() {
let (base_url, handle) = mock_server().await;
let tool = LoadTestingTool;
let mut params = HashMap::new();
params.insert(
"target_url".to_owned(),
serde_json::json!(format!("{base_url}/ok")),
);
let output = tool.execute(ToolInput { parameters: params }).await;
assert!(!output.success);
assert!(output.error.unwrap().contains("private"));
handle.abort();
}
}